1#[cfg(target_arch = "aarch64")]
19#[allow(unsafe_op_in_unsafe_fn)]
20mod neon {
21 use std::arch::aarch64::*;
22
23 #[target_feature(enable = "neon")]
25 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
26 let chunks = count / 16;
27 let remainder = count % 16;
28
29 for chunk in 0..chunks {
30 let base = chunk * 16;
31 let in_ptr = input.as_ptr().add(base);
32
33 let bytes = vld1q_u8(in_ptr);
35
36 let low8 = vget_low_u8(bytes);
38 let high8 = vget_high_u8(bytes);
39
40 let low16 = vmovl_u8(low8);
41 let high16 = vmovl_u8(high8);
42
43 let v0 = vmovl_u16(vget_low_u16(low16));
44 let v1 = vmovl_u16(vget_high_u16(low16));
45 let v2 = vmovl_u16(vget_low_u16(high16));
46 let v3 = vmovl_u16(vget_high_u16(high16));
47
48 let out_ptr = output.as_mut_ptr().add(base);
49 vst1q_u32(out_ptr, v0);
50 vst1q_u32(out_ptr.add(4), v1);
51 vst1q_u32(out_ptr.add(8), v2);
52 vst1q_u32(out_ptr.add(12), v3);
53 }
54
55 let base = chunks * 16;
57 for i in 0..remainder {
58 output[base + i] = input[base + i] as u32;
59 }
60 }
61
62 #[target_feature(enable = "neon")]
64 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
65 let chunks = count / 8;
66 let remainder = count % 8;
67
68 for chunk in 0..chunks {
69 let base = chunk * 8;
70 let in_ptr = input.as_ptr().add(base * 2) as *const u16;
71
72 let vals = vld1q_u16(in_ptr);
73 let low = vmovl_u16(vget_low_u16(vals));
74 let high = vmovl_u16(vget_high_u16(vals));
75
76 let out_ptr = output.as_mut_ptr().add(base);
77 vst1q_u32(out_ptr, low);
78 vst1q_u32(out_ptr.add(4), high);
79 }
80
81 let base = chunks * 8;
83 for i in 0..remainder {
84 let idx = (base + i) * 2;
85 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
86 }
87 }
88
89 #[target_feature(enable = "neon")]
91 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
92 let chunks = count / 4;
93 let remainder = count % 4;
94
95 let in_ptr = input.as_ptr() as *const u32;
96 let out_ptr = output.as_mut_ptr();
97
98 for chunk in 0..chunks {
99 let vals = vld1q_u32(in_ptr.add(chunk * 4));
100 vst1q_u32(out_ptr.add(chunk * 4), vals);
101 }
102
103 let base = chunks * 4;
105 for i in 0..remainder {
106 let idx = (base + i) * 4;
107 output[base + i] =
108 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
109 }
110 }
111
112 #[inline]
116 #[target_feature(enable = "neon")]
117 unsafe fn prefix_sum_4(v: uint32x4_t) -> uint32x4_t {
118 let shifted1 = vextq_u32(vdupq_n_u32(0), v, 3);
121 let sum1 = vaddq_u32(v, shifted1);
122
123 let shifted2 = vextq_u32(vdupq_n_u32(0), sum1, 2);
126 vaddq_u32(sum1, shifted2)
127 }
128
129 #[target_feature(enable = "neon")]
133 pub unsafe fn delta_decode(
134 output: &mut [u32],
135 deltas: &[u32],
136 first_doc_id: u32,
137 count: usize,
138 ) {
139 if count == 0 {
140 return;
141 }
142
143 output[0] = first_doc_id;
144 if count == 1 {
145 return;
146 }
147
148 let ones = vdupq_n_u32(1);
149 let mut carry = vdupq_n_u32(first_doc_id);
150
151 let full_groups = (count - 1) / 4;
152 let remainder = (count - 1) % 4;
153
154 for group in 0..full_groups {
155 let base = group * 4;
156
157 let d = vld1q_u32(deltas[base..].as_ptr());
159 let gaps = vaddq_u32(d, ones);
160
161 let prefix = prefix_sum_4(gaps);
163
164 let result = vaddq_u32(prefix, carry);
166
167 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
169
170 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
172 }
173
174 let base = full_groups * 4;
176 let mut scalar_carry = vgetq_lane_u32(carry, 0);
177 for j in 0..remainder {
178 scalar_carry = scalar_carry.wrapping_add(deltas[base + j]).wrapping_add(1);
179 output[base + j + 1] = scalar_carry;
180 }
181 }
182
183 #[target_feature(enable = "neon")]
185 pub unsafe fn add_one(values: &mut [u32], count: usize) {
186 let ones = vdupq_n_u32(1);
187 let chunks = count / 4;
188 let remainder = count % 4;
189
190 for chunk in 0..chunks {
191 let base = chunk * 4;
192 let ptr = values.as_mut_ptr().add(base);
193 let v = vld1q_u32(ptr);
194 let result = vaddq_u32(v, ones);
195 vst1q_u32(ptr, result);
196 }
197
198 let base = chunks * 4;
199 for i in 0..remainder {
200 values[base + i] += 1;
201 }
202 }
203
204 #[target_feature(enable = "neon")]
214 pub unsafe fn unpack_8bit_delta_decode_with_offset<const OFFSET: u32>(
215 input: &[u8],
216 output: &mut [u32],
217 first_value: u32,
218 count: usize,
219 ) {
220 output[0] = first_value;
221 if count <= 1 {
222 return;
223 }
224
225 let ones = vdupq_n_u32(OFFSET);
226 let mut carry = vdupq_n_u32(first_value);
227
228 let full_groups = (count - 1) / 4;
229 let remainder = (count - 1) % 4;
230
231 for group in 0..full_groups {
232 let base = group * 4;
233
234 let raw = std::ptr::read_unaligned(input.as_ptr().add(base) as *const u32);
236 let bytes = vreinterpret_u8_u32(vdup_n_u32(raw));
237 let u16s = vmovl_u8(bytes); let d = vmovl_u16(vget_low_u16(u16s)); let gaps = vaddq_u32(d, ones);
242
243 let prefix = prefix_sum_4(gaps);
245
246 let result = vaddq_u32(prefix, carry);
248
249 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
251
252 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
254 }
255
256 let base = full_groups * 4;
258 super::scalar::delta_decode_with_offset::<OFFSET, 1>(
259 &input[base..],
260 &mut output[base..],
261 vgetq_lane_u32(carry, 0),
262 remainder + 1,
263 );
264 }
265
266 #[target_feature(enable = "neon")]
276 pub unsafe fn unpack_16bit_delta_decode_with_offset<const OFFSET: u32>(
277 input: &[u8],
278 output: &mut [u32],
279 first_value: u32,
280 count: usize,
281 ) {
282 output[0] = first_value;
283 if count <= 1 {
284 return;
285 }
286
287 let ones = vdupq_n_u32(OFFSET);
288 let mut carry = vdupq_n_u32(first_value);
289
290 let full_groups = (count - 1) / 4;
291 let remainder = (count - 1) % 4;
292
293 for group in 0..full_groups {
294 let base = group * 4;
295 let in_ptr = input.as_ptr().add(base * 2) as *const u16;
296
297 let vals = vld1_u16(in_ptr);
299 let d = vmovl_u16(vals);
300
301 let gaps = vaddq_u32(d, ones);
303
304 let prefix = prefix_sum_4(gaps);
306
307 let result = vaddq_u32(prefix, carry);
309
310 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
312
313 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
315 }
316
317 let base = full_groups * 4;
319 super::scalar::delta_decode_with_offset::<OFFSET, 2>(
320 &input[base * 2..],
321 &mut output[base..],
322 vgetq_lane_u32(carry, 0),
323 remainder + 1,
324 );
325 }
326
327 #[target_feature(enable = "neon")]
330 pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
331 let len = a.len();
332 let chunks16 = len / 16;
333 let mut total = 0u32;
334
335 let mut i = 0;
338 while i < chunks16 {
339 let batch_end = (i + 31).min(chunks16);
340 let mut acc = vdupq_n_u8(0);
341 for j in i..batch_end {
342 let off = j * 16;
343 let va = vld1q_u8(a.as_ptr().add(off));
344 let vb = vld1q_u8(b.as_ptr().add(off));
345 let popcnt = vcntq_u8(veorq_u8(va, vb));
346 acc = vaddq_u8(acc, popcnt);
347 }
348 let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc)));
350 total += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
351 i = batch_end;
352 }
353
354 let base = chunks16 * 16;
356 for k in base..len {
357 total += (a[k] ^ b[k]).count_ones();
358 }
359
360 total
361 }
362
363 #[target_feature(enable = "neon")]
369 #[inline]
370 pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
371 let len = query.len();
372 let chunks16 = len / 16;
373 let mut total = [0u32; 4];
374
375 let mut i = 0;
376 while i < chunks16 {
377 let batch_end = (i + 31).min(chunks16);
378 let mut acc = [vdupq_n_u8(0); 4];
379 for j in i..batch_end {
380 let off = j * 16;
381 let vq = vld1q_u8(query.as_ptr().add(off));
382 for r in 0..4 {
383 let vr = vld1q_u8(rows[r].as_ptr().add(off));
384 acc[r] = vaddq_u8(acc[r], vcntq_u8(veorq_u8(vq, vr)));
385 }
386 }
387 for r in 0..4 {
388 let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc[r])));
389 total[r] += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
390 }
391 i = batch_end;
392 }
393
394 let base = chunks16 * 16;
397 if base < len {
398 let tail = &query[base..];
399 for r in 0..4 {
400 total[r] += super::hamming_distance_scalar(tail, &rows[r][base..]);
401 }
402 }
403
404 total
405 }
406
407 #[inline]
409 pub fn is_available() -> bool {
410 true
411 }
412}
413
414#[cfg(target_arch = "x86_64")]
419#[allow(unsafe_op_in_unsafe_fn)]
420mod sse {
421 use std::arch::x86_64::*;
422
423 #[target_feature(enable = "sse2", enable = "sse4.1")]
425 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
426 let chunks = count / 16;
427 let remainder = count % 16;
428
429 for chunk in 0..chunks {
430 let base = chunk * 16;
431 let in_ptr = input.as_ptr().add(base);
432
433 let bytes = _mm_loadu_si128(in_ptr as *const __m128i);
434
435 let v0 = _mm_cvtepu8_epi32(bytes);
437 let v1 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 4));
438 let v2 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 8));
439 let v3 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 12));
440
441 let out_ptr = output.as_mut_ptr().add(base);
442 _mm_storeu_si128(out_ptr as *mut __m128i, v0);
443 _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, v1);
444 _mm_storeu_si128(out_ptr.add(8) as *mut __m128i, v2);
445 _mm_storeu_si128(out_ptr.add(12) as *mut __m128i, v3);
446 }
447
448 let base = chunks * 16;
449 for i in 0..remainder {
450 output[base + i] = input[base + i] as u32;
451 }
452 }
453
454 #[target_feature(enable = "sse2", enable = "sse4.1")]
456 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
457 let chunks = count / 8;
458 let remainder = count % 8;
459
460 for chunk in 0..chunks {
461 let base = chunk * 8;
462 let in_ptr = input.as_ptr().add(base * 2);
463
464 let vals = _mm_loadu_si128(in_ptr as *const __m128i);
465 let low = _mm_cvtepu16_epi32(vals);
466 let high = _mm_cvtepu16_epi32(_mm_srli_si128(vals, 8));
467
468 let out_ptr = output.as_mut_ptr().add(base);
469 _mm_storeu_si128(out_ptr as *mut __m128i, low);
470 _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, high);
471 }
472
473 let base = chunks * 8;
474 for i in 0..remainder {
475 let idx = (base + i) * 2;
476 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
477 }
478 }
479
480 #[target_feature(enable = "sse2")]
482 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
483 let chunks = count / 4;
484 let remainder = count % 4;
485
486 let in_ptr = input.as_ptr() as *const __m128i;
487 let out_ptr = output.as_mut_ptr() as *mut __m128i;
488
489 for chunk in 0..chunks {
490 let vals = _mm_loadu_si128(in_ptr.add(chunk));
491 _mm_storeu_si128(out_ptr.add(chunk), vals);
492 }
493
494 let base = chunks * 4;
496 for i in 0..remainder {
497 let idx = (base + i) * 4;
498 output[base + i] =
499 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
500 }
501 }
502
503 #[inline]
507 #[target_feature(enable = "sse2")]
508 unsafe fn prefix_sum_4(v: __m128i) -> __m128i {
509 let shifted1 = _mm_slli_si128(v, 4);
512 let sum1 = _mm_add_epi32(v, shifted1);
513
514 let shifted2 = _mm_slli_si128(sum1, 8);
517 _mm_add_epi32(sum1, shifted2)
518 }
519
520 #[target_feature(enable = "sse2", enable = "sse4.1")]
522 pub unsafe fn delta_decode(
523 output: &mut [u32],
524 deltas: &[u32],
525 first_doc_id: u32,
526 count: usize,
527 ) {
528 if count == 0 {
529 return;
530 }
531
532 output[0] = first_doc_id;
533 if count == 1 {
534 return;
535 }
536
537 let ones = _mm_set1_epi32(1);
538 let mut carry = _mm_set1_epi32(first_doc_id as i32);
539
540 let full_groups = (count - 1) / 4;
541 let remainder = (count - 1) % 4;
542
543 for group in 0..full_groups {
544 let base = group * 4;
545
546 let d = _mm_loadu_si128(deltas[base..].as_ptr() as *const __m128i);
548 let gaps = _mm_add_epi32(d, ones);
549
550 let prefix = prefix_sum_4(gaps);
552
553 let result = _mm_add_epi32(prefix, carry);
555
556 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
558
559 carry = _mm_shuffle_epi32(result, 0xFF); }
562
563 let base = full_groups * 4;
565 let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
566 for j in 0..remainder {
567 scalar_carry = scalar_carry.wrapping_add(deltas[base + j]).wrapping_add(1);
568 output[base + j + 1] = scalar_carry;
569 }
570 }
571
572 #[target_feature(enable = "sse2")]
574 pub unsafe fn add_one(values: &mut [u32], count: usize) {
575 let ones = _mm_set1_epi32(1);
576 let chunks = count / 4;
577 let remainder = count % 4;
578
579 for chunk in 0..chunks {
580 let base = chunk * 4;
581 let ptr = values.as_mut_ptr().add(base) as *mut __m128i;
582 let v = _mm_loadu_si128(ptr);
583 let result = _mm_add_epi32(v, ones);
584 _mm_storeu_si128(ptr, result);
585 }
586
587 let base = chunks * 4;
588 for i in 0..remainder {
589 values[base + i] += 1;
590 }
591 }
592
593 #[target_feature(enable = "sse4.1")]
603 pub unsafe fn unpack_8bit_delta_decode_with_offset<const OFFSET: u32>(
604 input: &[u8],
605 output: &mut [u32],
606 first_value: u32,
607 count: usize,
608 ) {
609 output[0] = first_value;
610 if count <= 1 {
611 return;
612 }
613
614 let ones = _mm_set1_epi32(OFFSET as i32);
615 let mut carry = _mm_set1_epi32(first_value as i32);
616
617 let full_groups = (count - 1) / 4;
618 let remainder = (count - 1) % 4;
619
620 for group in 0..full_groups {
621 let base = group * 4;
622
623 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
625 input.as_ptr().add(base) as *const i32
626 ));
627 let d = _mm_cvtepu8_epi32(bytes);
628
629 let gaps = _mm_add_epi32(d, ones);
631
632 let prefix = prefix_sum_4(gaps);
634
635 let result = _mm_add_epi32(prefix, carry);
637
638 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
640
641 carry = _mm_shuffle_epi32(result, 0xFF);
643 }
644
645 let base = full_groups * 4;
647 super::scalar::delta_decode_with_offset::<OFFSET, 1>(
648 &input[base..],
649 &mut output[base..],
650 _mm_extract_epi32(carry, 0) as u32,
651 remainder + 1,
652 );
653 }
654
655 #[target_feature(enable = "sse4.1")]
665 pub unsafe fn unpack_16bit_delta_decode_with_offset<const OFFSET: u32>(
666 input: &[u8],
667 output: &mut [u32],
668 first_value: u32,
669 count: usize,
670 ) {
671 output[0] = first_value;
672 if count <= 1 {
673 return;
674 }
675
676 let ones = _mm_set1_epi32(OFFSET as i32);
677 let mut carry = _mm_set1_epi32(first_value as i32);
678
679 let full_groups = (count - 1) / 4;
680 let remainder = (count - 1) % 4;
681
682 for group in 0..full_groups {
683 let base = group * 4;
684 let in_ptr = input.as_ptr().add(base * 2);
685
686 let vals = _mm_loadl_epi64(in_ptr as *const __m128i); let d = _mm_cvtepu16_epi32(vals);
689
690 let gaps = _mm_add_epi32(d, ones);
692
693 let prefix = prefix_sum_4(gaps);
695
696 let result = _mm_add_epi32(prefix, carry);
698
699 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
701
702 carry = _mm_shuffle_epi32(result, 0xFF);
704 }
705
706 let base = full_groups * 4;
708 super::scalar::delta_decode_with_offset::<OFFSET, 2>(
709 &input[base * 2..],
710 &mut output[base..],
711 _mm_extract_epi32(carry, 0) as u32,
712 remainder + 1,
713 );
714 }
715
716 #[inline]
718 pub fn is_available() -> bool {
719 is_x86_feature_detected!("sse4.1")
720 }
721}
722
723#[cfg(target_arch = "x86_64")]
728#[allow(unsafe_op_in_unsafe_fn)]
729mod avx2 {
730 use std::arch::x86_64::*;
731
732 #[target_feature(enable = "avx2")]
734 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
735 let chunks = count / 32;
736 let remainder = count % 32;
737
738 for chunk in 0..chunks {
739 let base = chunk * 32;
740 let in_ptr = input.as_ptr().add(base);
741
742 let bytes_lo = _mm_loadu_si128(in_ptr as *const __m128i);
744 let bytes_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
745
746 let v0 = _mm256_cvtepu8_epi32(bytes_lo);
748 let v1 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_lo, 8));
749 let v2 = _mm256_cvtepu8_epi32(bytes_hi);
750 let v3 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_hi, 8));
751
752 let out_ptr = output.as_mut_ptr().add(base);
753 _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
754 _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
755 _mm256_storeu_si256(out_ptr.add(16) as *mut __m256i, v2);
756 _mm256_storeu_si256(out_ptr.add(24) as *mut __m256i, v3);
757 }
758
759 let base = chunks * 32;
761 for i in 0..remainder {
762 output[base + i] = input[base + i] as u32;
763 }
764 }
765
766 #[target_feature(enable = "avx2")]
768 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
769 let chunks = count / 16;
770 let remainder = count % 16;
771
772 for chunk in 0..chunks {
773 let base = chunk * 16;
774 let in_ptr = input.as_ptr().add(base * 2);
775
776 let vals_lo = _mm_loadu_si128(in_ptr as *const __m128i);
778 let vals_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
779
780 let v0 = _mm256_cvtepu16_epi32(vals_lo);
782 let v1 = _mm256_cvtepu16_epi32(vals_hi);
783
784 let out_ptr = output.as_mut_ptr().add(base);
785 _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
786 _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
787 }
788
789 let base = chunks * 16;
791 for i in 0..remainder {
792 let idx = (base + i) * 2;
793 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
794 }
795 }
796
797 #[target_feature(enable = "avx2")]
799 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
800 let chunks = count / 8;
801 let remainder = count % 8;
802
803 let in_ptr = input.as_ptr() as *const __m256i;
804 let out_ptr = output.as_mut_ptr() as *mut __m256i;
805
806 for chunk in 0..chunks {
807 let vals = _mm256_loadu_si256(in_ptr.add(chunk));
808 _mm256_storeu_si256(out_ptr.add(chunk), vals);
809 }
810
811 let base = chunks * 8;
813 for i in 0..remainder {
814 let idx = (base + i) * 4;
815 output[base + i] =
816 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
817 }
818 }
819
820 #[target_feature(enable = "avx2")]
822 pub unsafe fn add_one(values: &mut [u32], count: usize) {
823 let ones = _mm256_set1_epi32(1);
824 let chunks = count / 8;
825 let remainder = count % 8;
826
827 for chunk in 0..chunks {
828 let base = chunk * 8;
829 let ptr = values.as_mut_ptr().add(base) as *mut __m256i;
830 let v = _mm256_loadu_si256(ptr);
831 let result = _mm256_add_epi32(v, ones);
832 _mm256_storeu_si256(ptr, result);
833 }
834
835 let base = chunks * 8;
836 for i in 0..remainder {
837 values[base + i] += 1;
838 }
839 }
840
841 #[inline]
845 #[target_feature(enable = "avx2")]
846 unsafe fn prefix_sum_8(v: __m256i) -> __m256i {
847 let s1 = _mm256_slli_si256(v, 4);
849 let r1 = _mm256_add_epi32(v, s1);
850
851 let s2 = _mm256_slli_si256(r1, 8);
853 let r2 = _mm256_add_epi32(r1, s2);
854
855 let lo_sum = _mm256_shuffle_epi32(r2, 0xFF);
858 let carry = _mm256_permute2x128_si256(lo_sum, lo_sum, 0x00);
860 let carry_hi = _mm256_blend_epi32::<0xF0>(_mm256_setzero_si256(), carry);
862 _mm256_add_epi32(r2, carry_hi)
863 }
864
865 #[target_feature(enable = "avx2")]
875 pub unsafe fn unpack_8bit_delta_decode_with_offset<const OFFSET: u32>(
876 input: &[u8],
877 output: &mut [u32],
878 first_value: u32,
879 count: usize,
880 ) {
881 output[0] = first_value;
882 if count <= 1 {
883 return;
884 }
885
886 let ones = _mm256_set1_epi32(OFFSET as i32);
887 let mut carry = _mm256_set1_epi32(first_value as i32);
888 let broadcast_idx = _mm256_set1_epi32(7);
889
890 let full_groups = (count - 1) / 8;
891 let remainder = (count - 1) % 8;
892
893 for group in 0..full_groups {
894 let base = group * 8;
895
896 let bytes = _mm_loadl_epi64(input.as_ptr().add(base) as *const __m128i);
898 let d = _mm256_cvtepu8_epi32(bytes);
899
900 let gaps = _mm256_add_epi32(d, ones);
902
903 let prefix = prefix_sum_8(gaps);
905
906 let result = _mm256_add_epi32(prefix, carry);
908
909 _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
911
912 carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
914 }
915
916 let base = full_groups * 8;
918 super::scalar::delta_decode_with_offset::<OFFSET, 1>(
919 &input[base..],
920 &mut output[base..],
921 _mm256_extract_epi32::<0>(carry) as u32,
922 remainder + 1,
923 );
924 }
925
926 #[target_feature(enable = "avx2")]
936 pub unsafe fn unpack_16bit_delta_decode_with_offset<const OFFSET: u32>(
937 input: &[u8],
938 output: &mut [u32],
939 first_value: u32,
940 count: usize,
941 ) {
942 output[0] = first_value;
943 if count <= 1 {
944 return;
945 }
946
947 let ones = _mm256_set1_epi32(OFFSET as i32);
948 let mut carry = _mm256_set1_epi32(first_value as i32);
949 let broadcast_idx = _mm256_set1_epi32(7);
950
951 let full_groups = (count - 1) / 8;
952 let remainder = (count - 1) % 8;
953
954 for group in 0..full_groups {
955 let base = group * 8;
956 let in_ptr = input.as_ptr().add(base * 2);
957
958 let vals = _mm_loadu_si128(in_ptr as *const __m128i);
960 let d = _mm256_cvtepu16_epi32(vals);
961
962 let gaps = _mm256_add_epi32(d, ones);
964
965 let prefix = prefix_sum_8(gaps);
967
968 let result = _mm256_add_epi32(prefix, carry);
970
971 _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
973
974 carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
976 }
977
978 let base = full_groups * 8;
980 super::scalar::delta_decode_with_offset::<OFFSET, 2>(
981 &input[base * 2..],
982 &mut output[base..],
983 _mm256_extract_epi32::<0>(carry) as u32,
984 remainder + 1,
985 );
986 }
987
988 #[target_feature(enable = "avx2")]
991 pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
992 let len = a.len();
993 let chunks32 = len / 32;
994 let low_mask = _mm256_set1_epi8(0x0f);
995 let lookup = _mm256_setr_epi8(
997 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,
998 3, 3, 4,
999 );
1000 let mut total = 0u64;
1001
1002 let mut i = 0;
1003 while i < chunks32 {
1004 let batch_end = (i + 31).min(chunks32);
1007 let mut acc = _mm256_setzero_si256();
1008 for j in i..batch_end {
1009 let off = j * 32;
1010 let va = _mm256_loadu_si256(a.as_ptr().add(off) as *const __m256i);
1011 let vb = _mm256_loadu_si256(b.as_ptr().add(off) as *const __m256i);
1012 let xored = _mm256_xor_si256(va, vb);
1013 let lo = _mm256_and_si256(xored, low_mask);
1015 let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
1016 let popcnt = _mm256_add_epi8(
1017 _mm256_shuffle_epi8(lookup, lo),
1018 _mm256_shuffle_epi8(lookup, hi),
1019 );
1020 acc = _mm256_add_epi8(acc, popcnt);
1021 }
1022 let sad = _mm256_sad_epu8(acc, _mm256_setzero_si256());
1024 total += _mm256_extract_epi64(sad, 0) as u64
1025 + _mm256_extract_epi64(sad, 1) as u64
1026 + _mm256_extract_epi64(sad, 2) as u64
1027 + _mm256_extract_epi64(sad, 3) as u64;
1028 i = batch_end;
1029 }
1030
1031 let base = chunks32 * 32;
1033 for k in base..len {
1034 total += (a[k] ^ b[k]).count_ones() as u64;
1035 }
1036
1037 total as u32
1038 }
1039
1040 #[target_feature(enable = "avx2")]
1046 #[inline]
1047 pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
1048 let len = query.len();
1049 let chunks32 = len / 32;
1050 let low_mask = _mm256_set1_epi8(0x0f);
1051 let lookup = _mm256_setr_epi8(
1052 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,
1053 3, 3, 4,
1054 );
1055 let mut total = [0u64; 4];
1056
1057 let mut i = 0;
1058 while i < chunks32 {
1059 let batch_end = (i + 31).min(chunks32);
1060 let mut acc = [_mm256_setzero_si256(); 4];
1061 for j in i..batch_end {
1062 let off = j * 32;
1063 let vq = _mm256_loadu_si256(query.as_ptr().add(off) as *const __m256i);
1064 for r in 0..4 {
1065 let vr = _mm256_loadu_si256(rows[r].as_ptr().add(off) as *const __m256i);
1066 let xored = _mm256_xor_si256(vq, vr);
1067 let lo = _mm256_and_si256(xored, low_mask);
1068 let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
1069 acc[r] = _mm256_add_epi8(
1070 acc[r],
1071 _mm256_add_epi8(
1072 _mm256_shuffle_epi8(lookup, lo),
1073 _mm256_shuffle_epi8(lookup, hi),
1074 ),
1075 );
1076 }
1077 }
1078 for r in 0..4 {
1079 let sad = _mm256_sad_epu8(acc[r], _mm256_setzero_si256());
1080 total[r] += _mm256_extract_epi64(sad, 0) as u64
1081 + _mm256_extract_epi64(sad, 1) as u64
1082 + _mm256_extract_epi64(sad, 2) as u64
1083 + _mm256_extract_epi64(sad, 3) as u64;
1084 }
1085 i = batch_end;
1086 }
1087
1088 let base = chunks32 * 32;
1091 if base < len {
1092 let tail = &query[base..];
1093 for r in 0..4 {
1094 total[r] += u64::from(super::hamming_distance_scalar(tail, &rows[r][base..]));
1095 }
1096 }
1097
1098 [
1099 total[0] as u32,
1100 total[1] as u32,
1101 total[2] as u32,
1102 total[3] as u32,
1103 ]
1104 }
1105
1106 #[inline]
1108 pub fn is_available() -> bool {
1109 is_x86_feature_detected!("avx2")
1110 }
1111}
1112
1113mod scalar {
1118 #[inline]
1120 pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1121 for i in 0..count {
1122 output[i] = input[i] as u32;
1123 }
1124 }
1125
1126 #[inline]
1128 pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1129 for (i, out) in output.iter_mut().enumerate().take(count) {
1130 let idx = i * 2;
1131 *out = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
1132 }
1133 }
1134
1135 #[inline]
1137 pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1138 for (i, out) in output.iter_mut().enumerate().take(count) {
1139 let idx = i * 4;
1140 *out = u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
1141 }
1142 }
1143
1144 #[inline]
1146 pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_doc_id: u32, count: usize) {
1147 if count == 0 {
1148 return;
1149 }
1150
1151 output[0] = first_doc_id;
1152 let mut carry = first_doc_id;
1153
1154 for i in 0..count - 1 {
1155 carry = carry.wrapping_add(deltas[i]).wrapping_add(1);
1156 output[i + 1] = carry;
1157 }
1158 }
1159
1160 #[inline]
1162 pub fn add_one(values: &mut [u32], count: usize) {
1163 for val in values.iter_mut().take(count) {
1164 *val += 1;
1165 }
1166 }
1167
1168 #[inline]
1177 pub fn delta_decode_with_offset<const OFFSET: u32, const BYTES: usize>(
1178 input: &[u8],
1179 output: &mut [u32],
1180 first_value: u32,
1181 count: usize,
1182 ) {
1183 const {
1184 assert!(
1185 BYTES == 1 || BYTES == 2,
1186 "scalar delta decode supports 8/16-bit deltas"
1187 );
1188 }
1189 if count == 0 {
1190 return;
1191 }
1192 output[0] = first_value;
1193 let mut carry = first_value;
1194 for i in 0..count - 1 {
1195 let idx = i * BYTES;
1196 let delta = if BYTES == 1 {
1197 input[idx] as u32
1198 } else {
1199 u16::from_le_bytes([input[idx], input[idx + 1]]) as u32
1200 };
1201 carry = carry.wrapping_add(delta).wrapping_add(OFFSET);
1202 output[i + 1] = carry;
1203 }
1204 }
1205}
1206
1207#[inline]
1213pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1214 #[cfg(target_arch = "aarch64")]
1215 {
1216 if neon::is_available() {
1217 unsafe {
1218 neon::unpack_8bit(input, output, count);
1219 }
1220 return;
1221 }
1222 }
1223
1224 #[cfg(target_arch = "x86_64")]
1225 {
1226 if avx2::is_available() {
1228 unsafe {
1229 avx2::unpack_8bit(input, output, count);
1230 }
1231 return;
1232 }
1233 if sse::is_available() {
1234 unsafe {
1235 sse::unpack_8bit(input, output, count);
1236 }
1237 return;
1238 }
1239 }
1240
1241 scalar::unpack_8bit(input, output, count);
1242}
1243
1244#[inline]
1246pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1247 #[cfg(target_arch = "aarch64")]
1248 {
1249 if neon::is_available() {
1250 unsafe {
1251 neon::unpack_16bit(input, output, count);
1252 }
1253 return;
1254 }
1255 }
1256
1257 #[cfg(target_arch = "x86_64")]
1258 {
1259 if avx2::is_available() {
1261 unsafe {
1262 avx2::unpack_16bit(input, output, count);
1263 }
1264 return;
1265 }
1266 if sse::is_available() {
1267 unsafe {
1268 sse::unpack_16bit(input, output, count);
1269 }
1270 return;
1271 }
1272 }
1273
1274 scalar::unpack_16bit(input, output, count);
1275}
1276
1277#[inline]
1279pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1280 #[cfg(target_arch = "aarch64")]
1281 {
1282 if neon::is_available() {
1283 unsafe {
1284 neon::unpack_32bit(input, output, count);
1285 }
1286 return;
1287 }
1288 }
1289
1290 #[cfg(target_arch = "x86_64")]
1291 {
1292 if avx2::is_available() {
1294 unsafe {
1295 avx2::unpack_32bit(input, output, count);
1296 }
1297 return;
1298 }
1299 if sse::is_available() {
1300 unsafe {
1301 sse::unpack_32bit(input, output, count);
1302 }
1303 return;
1304 }
1305 }
1306
1307 scalar::unpack_32bit(input, output, count);
1308}
1309
1310#[inline]
1316pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_value: u32, count: usize) {
1317 #[cfg(target_arch = "aarch64")]
1318 {
1319 if neon::is_available() {
1320 unsafe {
1321 neon::delta_decode(output, deltas, first_value, count);
1322 }
1323 return;
1324 }
1325 }
1326
1327 #[cfg(target_arch = "x86_64")]
1328 {
1329 if sse::is_available() {
1330 unsafe {
1331 sse::delta_decode(output, deltas, first_value, count);
1332 }
1333 return;
1334 }
1335 }
1336
1337 scalar::delta_decode(output, deltas, first_value, count);
1338}
1339
1340#[inline]
1344pub fn add_one(values: &mut [u32], count: usize) {
1345 #[cfg(target_arch = "aarch64")]
1346 {
1347 if neon::is_available() {
1348 unsafe {
1349 neon::add_one(values, count);
1350 }
1351 return;
1352 }
1353 }
1354
1355 #[cfg(target_arch = "x86_64")]
1356 {
1357 if avx2::is_available() {
1359 unsafe {
1360 avx2::add_one(values, count);
1361 }
1362 return;
1363 }
1364 if sse::is_available() {
1365 unsafe {
1366 sse::add_one(values, count);
1367 }
1368 return;
1369 }
1370 }
1371
1372 scalar::add_one(values, count);
1373}
1374
1375#[inline]
1377pub fn bits_needed(val: u32) -> u8 {
1378 if val == 0 {
1379 0
1380 } else {
1381 32 - val.leading_zeros() as u8
1382 }
1383}
1384
1385#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1402#[repr(u8)]
1403pub enum RoundedBitWidth {
1404 Zero = 0,
1405 Bits8 = 8,
1406 Bits16 = 16,
1407 Bits32 = 32,
1408}
1409
1410impl RoundedBitWidth {
1411 #[inline]
1413 pub fn from_exact(bits: u8) -> Self {
1414 match bits {
1415 0 => RoundedBitWidth::Zero,
1416 1..=8 => RoundedBitWidth::Bits8,
1417 9..=16 => RoundedBitWidth::Bits16,
1418 _ => RoundedBitWidth::Bits32,
1419 }
1420 }
1421
1422 #[inline]
1426 pub fn try_from_u8(bits: u8) -> Option<Self> {
1427 match bits {
1428 0 => Some(RoundedBitWidth::Zero),
1429 8 => Some(RoundedBitWidth::Bits8),
1430 16 => Some(RoundedBitWidth::Bits16),
1431 32 => Some(RoundedBitWidth::Bits32),
1432 _ => None,
1433 }
1434 }
1435
1436 #[inline]
1440 pub fn from_u8(bits: u8) -> Self {
1441 Self::try_from_u8(bits).unwrap_or(RoundedBitWidth::Bits32)
1442 }
1443
1444 #[inline]
1446 pub fn bytes_per_value(self) -> usize {
1447 match self {
1448 RoundedBitWidth::Zero => 0,
1449 RoundedBitWidth::Bits8 => 1,
1450 RoundedBitWidth::Bits16 => 2,
1451 RoundedBitWidth::Bits32 => 4,
1452 }
1453 }
1454
1455 #[inline]
1457 pub fn as_u8(self) -> u8 {
1458 self as u8
1459 }
1460}
1461
1462#[inline]
1464pub fn round_bit_width(bits: u8) -> u8 {
1465 RoundedBitWidth::from_exact(bits).as_u8()
1466}
1467
1468#[inline]
1473pub fn pack_rounded(values: &[u32], bit_width: RoundedBitWidth, output: &mut [u8]) -> usize {
1474 let count = values.len();
1475 match bit_width {
1476 RoundedBitWidth::Zero => 0,
1477 RoundedBitWidth::Bits8 => {
1478 for (i, &v) in values.iter().enumerate() {
1479 output[i] = v as u8;
1480 }
1481 count
1482 }
1483 RoundedBitWidth::Bits16 => {
1484 for (i, &v) in values.iter().enumerate() {
1485 let bytes = (v as u16).to_le_bytes();
1486 output[i * 2] = bytes[0];
1487 output[i * 2 + 1] = bytes[1];
1488 }
1489 count * 2
1490 }
1491 RoundedBitWidth::Bits32 => {
1492 for (i, &v) in values.iter().enumerate() {
1493 let bytes = v.to_le_bytes();
1494 output[i * 4] = bytes[0];
1495 output[i * 4 + 1] = bytes[1];
1496 output[i * 4 + 2] = bytes[2];
1497 output[i * 4 + 3] = bytes[3];
1498 }
1499 count * 4
1500 }
1501 }
1502}
1503
1504#[inline]
1508pub fn unpack_rounded(input: &[u8], bit_width: RoundedBitWidth, output: &mut [u32], count: usize) {
1509 match bit_width {
1510 RoundedBitWidth::Zero => {
1511 for out in output.iter_mut().take(count) {
1512 *out = 0;
1513 }
1514 }
1515 RoundedBitWidth::Bits8 => unpack_8bit(input, output, count),
1516 RoundedBitWidth::Bits16 => unpack_16bit(input, output, count),
1517 RoundedBitWidth::Bits32 => unpack_32bit(input, output, count),
1518 }
1519}
1520
1521#[inline]
1524pub(crate) fn unpack_rounded_raw_delta_decode(
1525 input: &[u8],
1526 bit_width: RoundedBitWidth,
1527 output: &mut [u32],
1528 first_value: u32,
1529 count: usize,
1530) {
1531 match bit_width {
1532 RoundedBitWidth::Zero => output.iter_mut().take(count).for_each(|v| *v = first_value),
1533 RoundedBitWidth::Bits8 => {
1534 unpack_8bit_delta_decode_with_offset::<0>(input, output, first_value, count)
1535 }
1536 RoundedBitWidth::Bits16 => {
1537 unpack_16bit_delta_decode_with_offset::<0>(input, output, first_value, count)
1538 }
1539 RoundedBitWidth::Bits32 => {
1540 if count > 0 {
1541 output[0] = first_value;
1542 let mut carry = first_value;
1543 for i in 0..count - 1 {
1544 let offset = i * 4;
1545 let delta = u32::from_le_bytes(input[offset..offset + 4].try_into().unwrap());
1546 carry = carry.wrapping_add(delta);
1547 output[i + 1] = carry;
1548 }
1549 }
1550 }
1551 }
1552}
1553
1554#[inline]
1558pub fn unpack_rounded_delta_decode(
1559 input: &[u8],
1560 bit_width: RoundedBitWidth,
1561 output: &mut [u32],
1562 first_value: u32,
1563 count: usize,
1564) {
1565 match bit_width {
1566 RoundedBitWidth::Zero => {
1567 let mut val = first_value;
1569 for out in output.iter_mut().take(count) {
1570 *out = val;
1571 val = val.wrapping_add(1);
1572 }
1573 }
1574 RoundedBitWidth::Bits8 => unpack_8bit_delta_decode(input, output, first_value, count),
1575 RoundedBitWidth::Bits16 => unpack_16bit_delta_decode(input, output, first_value, count),
1576 RoundedBitWidth::Bits32 => {
1577 if count > 0 {
1579 output[0] = first_value;
1580 let mut carry = first_value;
1581 for i in 0..count - 1 {
1582 let idx = i * 4;
1583 let delta = u32::from_le_bytes([
1584 input[idx],
1585 input[idx + 1],
1586 input[idx + 2],
1587 input[idx + 3],
1588 ]);
1589 carry = carry.wrapping_add(delta).wrapping_add(1);
1590 output[i + 1] = carry;
1591 }
1592 }
1593 }
1594 }
1595}
1596
1597#[inline]
1606pub fn unpack_8bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1607 unpack_8bit_delta_decode_with_offset::<1>(input, output, first_value, count);
1608}
1609
1610#[inline]
1616pub(crate) fn unpack_8bit_delta_decode_with_offset<const OFFSET: u32>(
1617 input: &[u8],
1618 output: &mut [u32],
1619 first_value: u32,
1620 count: usize,
1621) {
1622 if count == 0 {
1623 return;
1624 }
1625 assert_delta_decode_bounds(input.len(), output.len(), count, 1);
1626
1627 output[0] = first_value;
1628 if count == 1 {
1629 return;
1630 }
1631
1632 #[cfg(target_arch = "aarch64")]
1633 {
1634 if neon::is_available() {
1635 unsafe {
1637 neon::unpack_8bit_delta_decode_with_offset::<OFFSET>(
1638 input,
1639 output,
1640 first_value,
1641 count,
1642 );
1643 }
1644 return;
1645 }
1646 }
1647
1648 #[cfg(target_arch = "x86_64")]
1649 {
1650 if avx2::is_available() {
1651 unsafe {
1653 avx2::unpack_8bit_delta_decode_with_offset::<OFFSET>(
1654 input,
1655 output,
1656 first_value,
1657 count,
1658 );
1659 }
1660 return;
1661 }
1662 if sse::is_available() {
1663 unsafe {
1665 sse::unpack_8bit_delta_decode_with_offset::<OFFSET>(
1666 input,
1667 output,
1668 first_value,
1669 count,
1670 );
1671 }
1672 return;
1673 }
1674 }
1675
1676 scalar::delta_decode_with_offset::<OFFSET, 1>(input, output, first_value, count);
1677}
1678
1679#[inline]
1681pub fn unpack_16bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1682 unpack_16bit_delta_decode_with_offset::<1>(input, output, first_value, count);
1683}
1684
1685#[inline]
1691pub(crate) fn unpack_16bit_delta_decode_with_offset<const OFFSET: u32>(
1692 input: &[u8],
1693 output: &mut [u32],
1694 first_value: u32,
1695 count: usize,
1696) {
1697 if count == 0 {
1698 return;
1699 }
1700 assert_delta_decode_bounds(input.len(), output.len(), count, 2);
1701
1702 output[0] = first_value;
1703 if count == 1 {
1704 return;
1705 }
1706
1707 #[cfg(target_arch = "aarch64")]
1708 {
1709 if neon::is_available() {
1710 unsafe {
1712 neon::unpack_16bit_delta_decode_with_offset::<OFFSET>(
1713 input,
1714 output,
1715 first_value,
1716 count,
1717 );
1718 }
1719 return;
1720 }
1721 }
1722
1723 #[cfg(target_arch = "x86_64")]
1724 {
1725 if avx2::is_available() {
1726 unsafe {
1728 avx2::unpack_16bit_delta_decode_with_offset::<OFFSET>(
1729 input,
1730 output,
1731 first_value,
1732 count,
1733 );
1734 }
1735 return;
1736 }
1737 if sse::is_available() {
1738 unsafe {
1740 sse::unpack_16bit_delta_decode_with_offset::<OFFSET>(
1741 input,
1742 output,
1743 first_value,
1744 count,
1745 );
1746 }
1747 return;
1748 }
1749 }
1750
1751 scalar::delta_decode_with_offset::<OFFSET, 2>(input, output, first_value, count);
1752}
1753
1754#[inline]
1758fn assert_delta_decode_bounds(input_len: usize, output_len: usize, count: usize, bytes: usize) {
1759 assert!(
1760 output_len >= count,
1761 "fused delta decode: output holds {output_len} values, block needs {count}"
1762 );
1763 let needed = (count - 1) * bytes;
1764 assert!(
1765 input_len >= needed,
1766 "fused delta decode: input holds {input_len} bytes, block needs {needed}"
1767 );
1768}
1769
1770#[inline]
1775pub fn unpack_delta_decode(
1776 input: &[u8],
1777 bit_width: u8,
1778 output: &mut [u32],
1779 first_value: u32,
1780 count: usize,
1781) {
1782 if count == 0 {
1783 return;
1784 }
1785
1786 output[0] = first_value;
1787 if count == 1 {
1788 return;
1789 }
1790
1791 match bit_width {
1793 0 => {
1794 let mut val = first_value;
1796 for item in output.iter_mut().take(count).skip(1) {
1797 val = val.wrapping_add(1);
1798 *item = val;
1799 }
1800 }
1801 8 => unpack_8bit_delta_decode(input, output, first_value, count),
1802 16 => unpack_16bit_delta_decode(input, output, first_value, count),
1803 32 => {
1804 let mut carry = first_value;
1806 for i in 0..count - 1 {
1807 let idx = i * 4;
1808 let delta = u32::from_le_bytes([
1809 input[idx],
1810 input[idx + 1],
1811 input[idx + 2],
1812 input[idx + 3],
1813 ]);
1814 carry = carry.wrapping_add(delta).wrapping_add(1);
1815 output[i + 1] = carry;
1816 }
1817 }
1818 _ => {
1819 let mask = (1u64 << bit_width) - 1;
1821 let bit_width_usize = bit_width as usize;
1822 let mut bit_pos = 0usize;
1823 let input_ptr = input.as_ptr();
1824 let mut carry = first_value;
1825
1826 for i in 0..count - 1 {
1827 let byte_idx = bit_pos >> 3;
1828 let bit_offset = bit_pos & 7;
1829
1830 let word = unsafe { (input_ptr.add(byte_idx) as *const u64).read_unaligned() };
1832 let delta = ((word >> bit_offset) & mask) as u32;
1833
1834 carry = carry.wrapping_add(delta).wrapping_add(1);
1835 output[i + 1] = carry;
1836 bit_pos += bit_width_usize;
1837 }
1838 }
1839 }
1840}
1841
1842#[inline]
1850pub fn dequantize_uint8(input: &[u8], output: &mut [f32], scale: f32, min_val: f32, count: usize) {
1851 #[cfg(target_arch = "aarch64")]
1852 {
1853 if neon::is_available() {
1854 unsafe {
1855 dequantize_uint8_neon(input, output, scale, min_val, count);
1856 }
1857 return;
1858 }
1859 }
1860
1861 #[cfg(target_arch = "x86_64")]
1862 {
1863 if sse::is_available() {
1864 unsafe {
1865 dequantize_uint8_sse(input, output, scale, min_val, count);
1866 }
1867 return;
1868 }
1869 }
1870
1871 for i in 0..count {
1873 output[i] = input[i] as f32 * scale + min_val;
1874 }
1875}
1876
1877#[cfg(target_arch = "aarch64")]
1878#[target_feature(enable = "neon")]
1879#[allow(unsafe_op_in_unsafe_fn)]
1880unsafe fn dequantize_uint8_neon(
1881 input: &[u8],
1882 output: &mut [f32],
1883 scale: f32,
1884 min_val: f32,
1885 count: usize,
1886) {
1887 use std::arch::aarch64::*;
1888
1889 let scale_v = vdupq_n_f32(scale);
1890 let min_v = vdupq_n_f32(min_val);
1891
1892 let chunks = count / 16;
1893 let remainder = count % 16;
1894
1895 for chunk in 0..chunks {
1896 let base = chunk * 16;
1897 let in_ptr = input.as_ptr().add(base);
1898
1899 let bytes = vld1q_u8(in_ptr);
1901
1902 let low8 = vget_low_u8(bytes);
1904 let high8 = vget_high_u8(bytes);
1905
1906 let low16 = vmovl_u8(low8);
1907 let high16 = vmovl_u8(high8);
1908
1909 let u32_0 = vmovl_u16(vget_low_u16(low16));
1911 let u32_1 = vmovl_u16(vget_high_u16(low16));
1912 let u32_2 = vmovl_u16(vget_low_u16(high16));
1913 let u32_3 = vmovl_u16(vget_high_u16(high16));
1914
1915 let f32_0 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_0), scale_v);
1917 let f32_1 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_1), scale_v);
1918 let f32_2 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_2), scale_v);
1919 let f32_3 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_3), scale_v);
1920
1921 let out_ptr = output.as_mut_ptr().add(base);
1922 vst1q_f32(out_ptr, f32_0);
1923 vst1q_f32(out_ptr.add(4), f32_1);
1924 vst1q_f32(out_ptr.add(8), f32_2);
1925 vst1q_f32(out_ptr.add(12), f32_3);
1926 }
1927
1928 let base = chunks * 16;
1930 for i in 0..remainder {
1931 output[base + i] = input[base + i] as f32 * scale + min_val;
1932 }
1933}
1934
1935#[cfg(target_arch = "x86_64")]
1936#[target_feature(enable = "sse2", enable = "sse4.1")]
1937#[allow(unsafe_op_in_unsafe_fn)]
1938unsafe fn dequantize_uint8_sse(
1939 input: &[u8],
1940 output: &mut [f32],
1941 scale: f32,
1942 min_val: f32,
1943 count: usize,
1944) {
1945 use std::arch::x86_64::*;
1946
1947 let scale_v = _mm_set1_ps(scale);
1948 let min_v = _mm_set1_ps(min_val);
1949
1950 let chunks = count / 4;
1951 let remainder = count % 4;
1952
1953 for chunk in 0..chunks {
1954 let base = chunk * 4;
1955
1956 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
1958 input.as_ptr().add(base) as *const i32
1959 ));
1960 let ints = _mm_cvtepu8_epi32(bytes);
1961 let floats = _mm_cvtepi32_ps(ints);
1962
1963 let scaled = _mm_add_ps(_mm_mul_ps(floats, scale_v), min_v);
1965
1966 _mm_storeu_ps(output.as_mut_ptr().add(base), scaled);
1967 }
1968
1969 let base = chunks * 4;
1971 for i in 0..remainder {
1972 output[base + i] = input[base + i] as f32 * scale + min_val;
1973 }
1974}
1975
1976#[inline]
2020fn dot_product_f32_scalar(a: &[f32], b: &[f32]) -> f32 {
2021 a.iter().zip(b).fold(0.0f32, |acc, (&x, &y)| {
2022 acc.algebraic_add(x.algebraic_mul(y))
2023 })
2024}
2025
2026#[inline]
2028fn fused_dot_norm_scalar(a: &[f32], b: &[f32]) -> (f32, f32) {
2029 a.iter()
2030 .zip(b)
2031 .fold((0.0f32, 0.0f32), |(dot, norm), (&x, &y)| {
2032 (
2033 dot.algebraic_add(x.algebraic_mul(y)),
2034 norm.algebraic_add(y.algebraic_mul(y)),
2035 )
2036 })
2037}
2038
2039#[inline]
2044pub fn squared_l2_f32(a: &[f32], b: &[f32]) -> f32 {
2045 a.iter().zip(b).fold(0.0f32, |acc, (&x, &y)| {
2046 let delta = x - y;
2047 acc.algebraic_add(delta.algebraic_mul(delta))
2048 })
2049}
2050
2051#[inline]
2053pub fn norm_squared_f32(v: &[f32]) -> f32 {
2054 v.iter()
2055 .fold(0.0f32, |acc, &x| acc.algebraic_add(x.algebraic_mul(x)))
2056}
2057
2058#[inline]
2060pub fn norm_f32(v: &[f32]) -> f32 {
2061 norm_squared_f32(v).sqrt()
2062}
2063
2064#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2070pub enum DenseF32Kernel {
2071 #[cfg(target_arch = "aarch64")]
2072 Neon,
2073 #[cfg(target_arch = "x86_64")]
2074 Avx512,
2075 #[cfg(target_arch = "x86_64")]
2076 Avx2Fma,
2077 #[cfg(target_arch = "x86_64")]
2078 Sse,
2079 Scalar,
2080}
2081
2082impl DenseF32Kernel {
2083 #[inline]
2085 pub fn resolve() -> Self {
2086 #[cfg(target_arch = "aarch64")]
2087 {
2088 if neon::is_available() {
2089 Self::Neon
2090 } else {
2091 Self::Scalar
2092 }
2093 }
2094 #[cfg(target_arch = "x86_64")]
2095 {
2096 if is_x86_feature_detected!("avx512f") {
2097 return Self::Avx512;
2098 }
2099 if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
2100 return Self::Avx2Fma;
2101 }
2102 if sse::is_available() {
2103 return Self::Sse;
2104 }
2105 Self::Scalar
2106 }
2107 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
2108 {
2109 Self::Scalar
2110 }
2111 }
2112
2113 #[inline]
2115 pub fn dot(self, a: &[f32], b: &[f32], count: usize) -> f32 {
2116 debug_assert!(count <= a.len() && count <= b.len());
2117 match self {
2118 #[cfg(target_arch = "aarch64")]
2119 Self::Neon => unsafe { dot_product_f32_neon(a, b, count) },
2120 #[cfg(target_arch = "x86_64")]
2121 Self::Avx512 => unsafe { dot_product_f32_avx512(a, b, count) },
2122 #[cfg(target_arch = "x86_64")]
2123 Self::Avx2Fma => unsafe { dot_product_f32_avx2(a, b, count) },
2124 #[cfg(target_arch = "x86_64")]
2125 Self::Sse => unsafe { dot_product_f32_sse(a, b, count) },
2126 Self::Scalar => dot_product_f32_scalar(&a[..count], &b[..count]),
2127 }
2128 }
2129
2130 #[inline]
2132 pub fn fused_dot_norm(self, a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2133 debug_assert!(count <= a.len() && count <= b.len());
2134 match self {
2135 #[cfg(target_arch = "aarch64")]
2136 Self::Neon => unsafe { fused_dot_norm_neon(a, b, count) },
2137 #[cfg(target_arch = "x86_64")]
2138 Self::Avx512 => unsafe { fused_dot_norm_avx512(a, b, count) },
2139 #[cfg(target_arch = "x86_64")]
2140 Self::Avx2Fma => unsafe { fused_dot_norm_avx2(a, b, count) },
2141 #[cfg(target_arch = "x86_64")]
2142 Self::Sse => unsafe { fused_dot_norm_sse(a, b, count) },
2143 Self::Scalar => fused_dot_norm_scalar(&a[..count], &b[..count]),
2144 }
2145 }
2146}
2147
2148#[inline]
2150pub fn dot_product_f32(a: &[f32], b: &[f32], count: usize) -> f32 {
2151 assert!(
2152 count <= a.len() && count <= b.len(),
2153 "dot_product_f32 count {count} exceeds input lengths ({}, {})",
2154 a.len(),
2155 b.len()
2156 );
2157 DenseF32Kernel::resolve().dot(a, b, count)
2158}
2159
2160#[cfg(target_arch = "aarch64")]
2161#[target_feature(enable = "neon")]
2162#[allow(unsafe_op_in_unsafe_fn)]
2163unsafe fn dot_product_f32_neon(a: &[f32], b: &[f32], count: usize) -> f32 {
2164 use std::arch::aarch64::*;
2165
2166 let chunks16 = count / 16;
2167 let remainder = count % 16;
2168
2169 let mut acc0 = vdupq_n_f32(0.0);
2170 let mut acc1 = vdupq_n_f32(0.0);
2171 let mut acc2 = vdupq_n_f32(0.0);
2172 let mut acc3 = vdupq_n_f32(0.0);
2173
2174 for c in 0..chunks16 {
2175 let base = c * 16;
2176 acc0 = vfmaq_f32(
2177 acc0,
2178 vld1q_f32(a.as_ptr().add(base)),
2179 vld1q_f32(b.as_ptr().add(base)),
2180 );
2181 acc1 = vfmaq_f32(
2182 acc1,
2183 vld1q_f32(a.as_ptr().add(base + 4)),
2184 vld1q_f32(b.as_ptr().add(base + 4)),
2185 );
2186 acc2 = vfmaq_f32(
2187 acc2,
2188 vld1q_f32(a.as_ptr().add(base + 8)),
2189 vld1q_f32(b.as_ptr().add(base + 8)),
2190 );
2191 acc3 = vfmaq_f32(
2192 acc3,
2193 vld1q_f32(a.as_ptr().add(base + 12)),
2194 vld1q_f32(b.as_ptr().add(base + 12)),
2195 );
2196 }
2197
2198 let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
2199 let mut sum = vaddvq_f32(acc);
2200
2201 let mut base = chunks16 * 16;
2204 if remainder >= 4 && remainder != 8 {
2208 let mut tail = vdupq_n_f32(0.0);
2209 while base + 4 <= count {
2210 tail = vfmaq_f32(
2211 tail,
2212 vld1q_f32(a.as_ptr().add(base)),
2213 vld1q_f32(b.as_ptr().add(base)),
2214 );
2215 base += 4;
2216 }
2217 sum += vaddvq_f32(tail);
2218 }
2219 for i in base..count {
2220 sum = sum.algebraic_add(a[i].algebraic_mul(b[i]));
2221 }
2222
2223 sum
2224}
2225
2226#[cfg(target_arch = "x86_64")]
2227#[target_feature(enable = "avx2", enable = "fma")]
2228#[allow(unsafe_op_in_unsafe_fn)]
2229unsafe fn dot_product_f32_avx2(a: &[f32], b: &[f32], count: usize) -> f32 {
2230 use std::arch::x86_64::*;
2231
2232 let chunks32 = count / 32;
2233 let remainder = count % 32;
2234
2235 let mut acc0 = _mm256_setzero_ps();
2236 let mut acc1 = _mm256_setzero_ps();
2237 let mut acc2 = _mm256_setzero_ps();
2238 let mut acc3 = _mm256_setzero_ps();
2239
2240 for c in 0..chunks32 {
2241 let base = c * 32;
2242 acc0 = _mm256_fmadd_ps(
2243 _mm256_loadu_ps(a.as_ptr().add(base)),
2244 _mm256_loadu_ps(b.as_ptr().add(base)),
2245 acc0,
2246 );
2247 acc1 = _mm256_fmadd_ps(
2248 _mm256_loadu_ps(a.as_ptr().add(base + 8)),
2249 _mm256_loadu_ps(b.as_ptr().add(base + 8)),
2250 acc1,
2251 );
2252 acc2 = _mm256_fmadd_ps(
2253 _mm256_loadu_ps(a.as_ptr().add(base + 16)),
2254 _mm256_loadu_ps(b.as_ptr().add(base + 16)),
2255 acc2,
2256 );
2257 acc3 = _mm256_fmadd_ps(
2258 _mm256_loadu_ps(a.as_ptr().add(base + 24)),
2259 _mm256_loadu_ps(b.as_ptr().add(base + 24)),
2260 acc3,
2261 );
2262 }
2263
2264 let acc = _mm256_add_ps(_mm256_add_ps(acc0, acc1), _mm256_add_ps(acc2, acc3));
2265
2266 let hi = _mm256_extractf128_ps(acc, 1);
2268 let lo = _mm256_castps256_ps128(acc);
2269 let sum128 = _mm_add_ps(lo, hi);
2270 let shuf = _mm_shuffle_ps(sum128, sum128, 0b10_11_00_01);
2271 let sums = _mm_add_ps(sum128, shuf);
2272 let shuf2 = _mm_movehl_ps(sums, sums);
2273 let final_sum = _mm_add_ss(sums, shuf2);
2274
2275 let mut sum = _mm_cvtss_f32(final_sum);
2276
2277 let mut base = chunks32 * 32;
2280 if remainder >= 8 {
2281 let mut tail = _mm256_setzero_ps();
2282 while base + 8 <= count {
2283 tail = _mm256_fmadd_ps(
2284 _mm256_loadu_ps(a.as_ptr().add(base)),
2285 _mm256_loadu_ps(b.as_ptr().add(base)),
2286 tail,
2287 );
2288 base += 8;
2289 }
2290 let hi = _mm256_extractf128_ps(tail, 1);
2291 let lo = _mm256_castps256_ps128(tail);
2292 let sum128 = _mm_add_ps(lo, hi);
2293 let shuf = _mm_shuffle_ps(sum128, sum128, 0b10_11_00_01);
2294 let sums = _mm_add_ps(sum128, shuf);
2295 let shuf2 = _mm_movehl_ps(sums, sums);
2296 sum += _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
2297 }
2298 for i in base..count {
2299 sum = sum.algebraic_add(a[i].algebraic_mul(b[i]));
2300 }
2301
2302 sum
2303}
2304
2305#[cfg(target_arch = "x86_64")]
2306#[target_feature(enable = "sse")]
2307#[allow(unsafe_op_in_unsafe_fn)]
2308unsafe fn dot_product_f32_sse(a: &[f32], b: &[f32], count: usize) -> f32 {
2309 use std::arch::x86_64::*;
2310
2311 let chunks = count / 4;
2312 let remainder = count % 4;
2313
2314 let mut acc = _mm_setzero_ps();
2315
2316 for chunk in 0..chunks {
2317 let base = chunk * 4;
2318 let va = _mm_loadu_ps(a.as_ptr().add(base));
2319 let vb = _mm_loadu_ps(b.as_ptr().add(base));
2320 acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
2321 }
2322
2323 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);
2330
2331 let base = chunks * 4;
2333 for i in 0..remainder {
2334 sum = sum.algebraic_add(a[base + i].algebraic_mul(b[base + i]));
2335 }
2336
2337 sum
2338}
2339
2340#[cfg(target_arch = "x86_64")]
2341#[target_feature(enable = "avx512f")]
2342#[allow(unsafe_op_in_unsafe_fn)]
2343unsafe fn dot_product_f32_avx512(a: &[f32], b: &[f32], count: usize) -> f32 {
2344 use std::arch::x86_64::*;
2345
2346 let chunks64 = count / 64;
2347 let remainder = count % 64;
2348
2349 let mut acc0 = _mm512_setzero_ps();
2350 let mut acc1 = _mm512_setzero_ps();
2351 let mut acc2 = _mm512_setzero_ps();
2352 let mut acc3 = _mm512_setzero_ps();
2353
2354 for c in 0..chunks64 {
2355 let base = c * 64;
2356 acc0 = _mm512_fmadd_ps(
2357 _mm512_loadu_ps(a.as_ptr().add(base)),
2358 _mm512_loadu_ps(b.as_ptr().add(base)),
2359 acc0,
2360 );
2361 acc1 = _mm512_fmadd_ps(
2362 _mm512_loadu_ps(a.as_ptr().add(base + 16)),
2363 _mm512_loadu_ps(b.as_ptr().add(base + 16)),
2364 acc1,
2365 );
2366 acc2 = _mm512_fmadd_ps(
2367 _mm512_loadu_ps(a.as_ptr().add(base + 32)),
2368 _mm512_loadu_ps(b.as_ptr().add(base + 32)),
2369 acc2,
2370 );
2371 acc3 = _mm512_fmadd_ps(
2372 _mm512_loadu_ps(a.as_ptr().add(base + 48)),
2373 _mm512_loadu_ps(b.as_ptr().add(base + 48)),
2374 acc3,
2375 );
2376 }
2377
2378 let acc = _mm512_add_ps(_mm512_add_ps(acc0, acc1), _mm512_add_ps(acc2, acc3));
2379 let mut sum = _mm512_reduce_add_ps(acc);
2380
2381 let mut base = chunks64 * 64;
2384 if remainder >= 16 {
2385 let mut tail = _mm512_setzero_ps();
2386 while base + 16 <= count {
2387 tail = _mm512_fmadd_ps(
2388 _mm512_loadu_ps(a.as_ptr().add(base)),
2389 _mm512_loadu_ps(b.as_ptr().add(base)),
2390 tail,
2391 );
2392 base += 16;
2393 }
2394 sum += _mm512_reduce_add_ps(tail);
2395 }
2396 for i in base..count {
2397 sum = sum.algebraic_add(a[i].algebraic_mul(b[i]));
2398 }
2399
2400 sum
2401}
2402
2403#[cfg(target_arch = "x86_64")]
2404#[target_feature(enable = "avx512f")]
2405#[allow(unsafe_op_in_unsafe_fn)]
2406unsafe fn fused_dot_norm_avx512(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2407 use std::arch::x86_64::*;
2408
2409 let chunks64 = count / 64;
2410 let remainder = count % 64;
2411
2412 let mut d0 = _mm512_setzero_ps();
2413 let mut d1 = _mm512_setzero_ps();
2414 let mut d2 = _mm512_setzero_ps();
2415 let mut d3 = _mm512_setzero_ps();
2416 let mut n0 = _mm512_setzero_ps();
2417 let mut n1 = _mm512_setzero_ps();
2418 let mut n2 = _mm512_setzero_ps();
2419 let mut n3 = _mm512_setzero_ps();
2420
2421 for c in 0..chunks64 {
2422 let base = c * 64;
2423 let vb0 = _mm512_loadu_ps(b.as_ptr().add(base));
2424 d0 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2425 n0 = _mm512_fmadd_ps(vb0, vb0, n0);
2426 let vb1 = _mm512_loadu_ps(b.as_ptr().add(base + 16));
2427 d1 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 16)), vb1, d1);
2428 n1 = _mm512_fmadd_ps(vb1, vb1, n1);
2429 let vb2 = _mm512_loadu_ps(b.as_ptr().add(base + 32));
2430 d2 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 32)), vb2, d2);
2431 n2 = _mm512_fmadd_ps(vb2, vb2, n2);
2432 let vb3 = _mm512_loadu_ps(b.as_ptr().add(base + 48));
2433 d3 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 48)), vb3, d3);
2434 n3 = _mm512_fmadd_ps(vb3, vb3, n3);
2435 }
2436
2437 let acc_dot = _mm512_add_ps(_mm512_add_ps(d0, d1), _mm512_add_ps(d2, d3));
2438 let acc_norm = _mm512_add_ps(_mm512_add_ps(n0, n1), _mm512_add_ps(n2, n3));
2439 let mut dot = _mm512_reduce_add_ps(acc_dot);
2440 let mut norm = _mm512_reduce_add_ps(acc_norm);
2441
2442 let mut base = chunks64 * 64;
2443 if remainder >= 16 {
2444 let mut tail_dot = _mm512_setzero_ps();
2445 let mut tail_norm = _mm512_setzero_ps();
2446 while base + 16 <= count {
2447 let vb = _mm512_loadu_ps(b.as_ptr().add(base));
2448 tail_dot = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base)), vb, tail_dot);
2449 tail_norm = _mm512_fmadd_ps(vb, vb, tail_norm);
2450 base += 16;
2451 }
2452 dot += _mm512_reduce_add_ps(tail_dot);
2453 norm += _mm512_reduce_add_ps(tail_norm);
2454 }
2455 for i in base..count {
2456 dot = dot.algebraic_add(a[i].algebraic_mul(b[i]));
2457 norm = norm.algebraic_add(b[i].algebraic_mul(b[i]));
2458 }
2459
2460 (dot, norm)
2461}
2462
2463#[inline]
2472fn fused_dot_norm(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2473 DenseF32Kernel::resolve().fused_dot_norm(a, b, count)
2474}
2475
2476#[cfg(target_arch = "aarch64")]
2477#[target_feature(enable = "neon")]
2478#[allow(unsafe_op_in_unsafe_fn)]
2479unsafe fn fused_dot_norm_neon(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2480 use std::arch::aarch64::*;
2481
2482 let chunks16 = count / 16;
2483 let remainder = count % 16;
2484
2485 let mut d0 = vdupq_n_f32(0.0);
2486 let mut d1 = vdupq_n_f32(0.0);
2487 let mut d2 = vdupq_n_f32(0.0);
2488 let mut d3 = vdupq_n_f32(0.0);
2489 let mut n0 = vdupq_n_f32(0.0);
2490 let mut n1 = vdupq_n_f32(0.0);
2491 let mut n2 = vdupq_n_f32(0.0);
2492 let mut n3 = vdupq_n_f32(0.0);
2493
2494 for c in 0..chunks16 {
2495 let base = c * 16;
2496 let va0 = vld1q_f32(a.as_ptr().add(base));
2497 let vb0 = vld1q_f32(b.as_ptr().add(base));
2498 d0 = vfmaq_f32(d0, va0, vb0);
2499 n0 = vfmaq_f32(n0, vb0, vb0);
2500 let va1 = vld1q_f32(a.as_ptr().add(base + 4));
2501 let vb1 = vld1q_f32(b.as_ptr().add(base + 4));
2502 d1 = vfmaq_f32(d1, va1, vb1);
2503 n1 = vfmaq_f32(n1, vb1, vb1);
2504 let va2 = vld1q_f32(a.as_ptr().add(base + 8));
2505 let vb2 = vld1q_f32(b.as_ptr().add(base + 8));
2506 d2 = vfmaq_f32(d2, va2, vb2);
2507 n2 = vfmaq_f32(n2, vb2, vb2);
2508 let va3 = vld1q_f32(a.as_ptr().add(base + 12));
2509 let vb3 = vld1q_f32(b.as_ptr().add(base + 12));
2510 d3 = vfmaq_f32(d3, va3, vb3);
2511 n3 = vfmaq_f32(n3, vb3, vb3);
2512 }
2513
2514 let acc_dot = vaddq_f32(vaddq_f32(d0, d1), vaddq_f32(d2, d3));
2515 let acc_norm = vaddq_f32(vaddq_f32(n0, n1), vaddq_f32(n2, n3));
2516 let mut dot = vaddvq_f32(acc_dot);
2517 let mut norm = vaddvq_f32(acc_norm);
2518
2519 let mut base = chunks16 * 16;
2520 if remainder >= 4 {
2521 let mut tail_dot = vdupq_n_f32(0.0);
2522 let mut tail_norm = vdupq_n_f32(0.0);
2523 while base + 4 <= count {
2524 let vb = vld1q_f32(b.as_ptr().add(base));
2525 tail_dot = vfmaq_f32(tail_dot, vld1q_f32(a.as_ptr().add(base)), vb);
2526 tail_norm = vfmaq_f32(tail_norm, vb, vb);
2527 base += 4;
2528 }
2529 dot += vaddvq_f32(tail_dot);
2530 norm += vaddvq_f32(tail_norm);
2531 }
2532 for i in base..count {
2533 dot = dot.algebraic_add(a[i].algebraic_mul(b[i]));
2534 norm = norm.algebraic_add(b[i].algebraic_mul(b[i]));
2535 }
2536
2537 (dot, norm)
2538}
2539
2540#[cfg(target_arch = "x86_64")]
2541#[target_feature(enable = "avx2", enable = "fma")]
2542#[allow(unsafe_op_in_unsafe_fn)]
2543unsafe fn fused_dot_norm_avx2(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2544 use std::arch::x86_64::*;
2545
2546 let chunks32 = count / 32;
2547 let remainder = count % 32;
2548
2549 let mut d0 = _mm256_setzero_ps();
2550 let mut d1 = _mm256_setzero_ps();
2551 let mut d2 = _mm256_setzero_ps();
2552 let mut d3 = _mm256_setzero_ps();
2553 let mut n0 = _mm256_setzero_ps();
2554 let mut n1 = _mm256_setzero_ps();
2555 let mut n2 = _mm256_setzero_ps();
2556 let mut n3 = _mm256_setzero_ps();
2557
2558 for c in 0..chunks32 {
2559 let base = c * 32;
2560 let vb0 = _mm256_loadu_ps(b.as_ptr().add(base));
2561 d0 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2562 n0 = _mm256_fmadd_ps(vb0, vb0, n0);
2563 let vb1 = _mm256_loadu_ps(b.as_ptr().add(base + 8));
2564 d1 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 8)), vb1, d1);
2565 n1 = _mm256_fmadd_ps(vb1, vb1, n1);
2566 let vb2 = _mm256_loadu_ps(b.as_ptr().add(base + 16));
2567 d2 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 16)), vb2, d2);
2568 n2 = _mm256_fmadd_ps(vb2, vb2, n2);
2569 let vb3 = _mm256_loadu_ps(b.as_ptr().add(base + 24));
2570 d3 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 24)), vb3, d3);
2571 n3 = _mm256_fmadd_ps(vb3, vb3, n3);
2572 }
2573
2574 let acc_dot = _mm256_add_ps(_mm256_add_ps(d0, d1), _mm256_add_ps(d2, d3));
2575 let acc_norm = _mm256_add_ps(_mm256_add_ps(n0, n1), _mm256_add_ps(n2, n3));
2576
2577 let hi_d = _mm256_extractf128_ps(acc_dot, 1);
2579 let lo_d = _mm256_castps256_ps128(acc_dot);
2580 let sum_d = _mm_add_ps(lo_d, hi_d);
2581 let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
2582 let sums_d = _mm_add_ps(sum_d, shuf_d);
2583 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2584 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2585
2586 let hi_n = _mm256_extractf128_ps(acc_norm, 1);
2587 let lo_n = _mm256_castps256_ps128(acc_norm);
2588 let sum_n = _mm_add_ps(lo_n, hi_n);
2589 let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
2590 let sums_n = _mm_add_ps(sum_n, shuf_n);
2591 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2592 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2593
2594 let mut base = chunks32 * 32;
2595 if remainder >= 8 {
2596 let mut tail_dot = _mm256_setzero_ps();
2597 let mut tail_norm = _mm256_setzero_ps();
2598 while base + 8 <= count {
2599 let vb = _mm256_loadu_ps(b.as_ptr().add(base));
2600 tail_dot = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base)), vb, tail_dot);
2601 tail_norm = _mm256_fmadd_ps(vb, vb, tail_norm);
2602 base += 8;
2603 }
2604 let reduce = |v: __m256| -> f32 {
2605 let hi = _mm256_extractf128_ps(v, 1);
2606 let lo = _mm256_castps256_ps128(v);
2607 let sum128 = _mm_add_ps(lo, hi);
2608 let shuf = _mm_shuffle_ps(sum128, sum128, 0b10_11_00_01);
2609 let sums = _mm_add_ps(sum128, shuf);
2610 let shuf2 = _mm_movehl_ps(sums, sums);
2611 _mm_cvtss_f32(_mm_add_ss(sums, shuf2))
2612 };
2613 dot += reduce(tail_dot);
2614 norm += reduce(tail_norm);
2615 }
2616 for i in base..count {
2617 dot = dot.algebraic_add(a[i].algebraic_mul(b[i]));
2618 norm = norm.algebraic_add(b[i].algebraic_mul(b[i]));
2619 }
2620
2621 (dot, norm)
2622}
2623
2624#[cfg(target_arch = "x86_64")]
2625#[target_feature(enable = "sse")]
2626#[allow(unsafe_op_in_unsafe_fn)]
2627unsafe fn fused_dot_norm_sse(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2628 use std::arch::x86_64::*;
2629
2630 let chunks = count / 4;
2631 let remainder = count % 4;
2632
2633 let mut acc_dot = _mm_setzero_ps();
2634 let mut acc_norm = _mm_setzero_ps();
2635
2636 for chunk in 0..chunks {
2637 let base = chunk * 4;
2638 let va = _mm_loadu_ps(a.as_ptr().add(base));
2639 let vb = _mm_loadu_ps(b.as_ptr().add(base));
2640 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2641 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2642 }
2643
2644 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2646 let sums_d = _mm_add_ps(acc_dot, shuf_d);
2647 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2648 let final_d = _mm_add_ss(sums_d, shuf2_d);
2649 let mut dot = _mm_cvtss_f32(final_d);
2650
2651 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2652 let sums_n = _mm_add_ps(acc_norm, shuf_n);
2653 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2654 let final_n = _mm_add_ss(sums_n, shuf2_n);
2655 let mut norm = _mm_cvtss_f32(final_n);
2656
2657 let base = chunks * 4;
2658 for i in 0..remainder {
2659 dot = dot.algebraic_add(a[base + i].algebraic_mul(b[base + i]));
2660 norm = norm.algebraic_add(b[base + i].algebraic_mul(b[base + i]));
2661 }
2662
2663 (dot, norm)
2664}
2665
2666#[inline]
2672pub fn fast_inv_sqrt(x: f32) -> f32 {
2673 let half = 0.5 * x;
2674 let i = 0x5F37_5A86_u32.wrapping_sub(x.to_bits() >> 1);
2675 let y = f32::from_bits(i);
2676 let y = y * (1.5 - half * y * y); y * (1.5 - half * y * y) }
2679
2680#[inline]
2691pub fn batch_cosine_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
2692 let n = scores.len();
2693 let required = n
2694 .checked_mul(dim)
2695 .expect("batch cosine vector length overflow");
2696 assert_eq!(query.len(), dim, "batch cosine query dimension mismatch");
2697 assert!(
2698 vectors.len() >= required,
2699 "batch cosine vectors are truncated: need {required}, got {}",
2700 vectors.len()
2701 );
2702
2703 if dim == 0 || n == 0 {
2704 return;
2705 }
2706
2707 let norm_q_sq = dot_product_f32(query, query, dim);
2709 if norm_q_sq < f32::EPSILON {
2710 for s in scores.iter_mut() {
2711 *s = 0.0;
2712 }
2713 return;
2714 }
2715 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
2716
2717 for i in 0..n {
2718 let vec = &vectors[i * dim..(i + 1) * dim];
2719 let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
2720 if norm_v_sq < f32::EPSILON {
2721 scores[i] = 0.0;
2722 } else {
2723 scores[i] = dot * inv_norm_q * fast_inv_sqrt(norm_v_sq);
2724 }
2725 }
2726}
2727
2728#[inline]
2734pub fn f32_to_f16(value: f32) -> u16 {
2735 let bits = value.to_bits();
2736 let sign = (bits >> 16) & 0x8000;
2737 let exp = ((bits >> 23) & 0xFF) as i32;
2738 let mantissa = bits & 0x7F_FFFF;
2739
2740 if exp == 255 {
2741 return (sign | 0x7C00 | ((mantissa >> 13) & 0x3FF)) as u16;
2743 }
2744
2745 let exp16 = exp - 127 + 15;
2746
2747 if exp16 >= 31 {
2748 return (sign | 0x7C00) as u16; }
2750
2751 if exp16 <= 0 {
2752 if exp16 < -10 {
2753 return sign as u16; }
2755 let shift = (1 - exp16) as u32;
2756 let m = (mantissa | 0x80_0000) >> shift;
2757 let round_bit = (m >> 12) & 1;
2759 let sticky = m & 0xFFF;
2760 let m13 = m >> 13;
2761 let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2762 return (sign | rounded) as u16;
2763 }
2764
2765 let round_bit = (mantissa >> 12) & 1;
2767 let sticky = mantissa & 0xFFF;
2768 let m13 = mantissa >> 13;
2769 let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2770 if rounded > 0x3FF {
2772 let exp16_inc = exp16 as u32 + 1;
2773 if exp16_inc >= 31 {
2774 return (sign | 0x7C00) as u16; }
2776 (sign | (exp16_inc << 10)) as u16
2777 } else {
2778 (sign | ((exp16 as u32) << 10) | rounded) as u16
2779 }
2780}
2781
2782#[inline]
2784pub fn f16_to_f32(half: u16) -> f32 {
2785 let sign = ((half & 0x8000) as u32) << 16;
2786 let exp = ((half >> 10) & 0x1F) as u32;
2787 let mantissa = (half & 0x3FF) as u32;
2788
2789 if exp == 0 {
2790 if mantissa == 0 {
2791 return f32::from_bits(sign);
2792 }
2793 let mut e = 0u32;
2795 let mut m = mantissa;
2796 while (m & 0x400) == 0 {
2797 m <<= 1;
2798 e += 1;
2799 }
2800 return f32::from_bits(sign | ((127 - 15 + 1 - e) << 23) | ((m & 0x3FF) << 13));
2801 }
2802
2803 if exp == 31 {
2804 return f32::from_bits(sign | 0x7F80_0000 | (mantissa << 13));
2805 }
2806
2807 f32::from_bits(sign | ((exp + 127 - 15) << 23) | (mantissa << 13))
2808}
2809
2810const U8_SCALE: f32 = 127.5;
2815const U8_INV_SCALE: f32 = 1.0 / 127.5;
2816
2817#[inline]
2819pub fn f32_to_u8_saturating(value: f32) -> u8 {
2820 ((value.clamp(-1.0, 1.0) + 1.0) * U8_SCALE) as u8
2821}
2822
2823#[inline]
2825pub fn u8_to_f32(byte: u8) -> f32 {
2826 byte as f32 * U8_INV_SCALE - 1.0
2827}
2828
2829pub fn batch_f32_to_f16(src: &[f32], dst: &mut [u16]) {
2835 debug_assert_eq!(src.len(), dst.len());
2836 for (s, d) in src.iter().zip(dst.iter_mut()) {
2837 *d = f32_to_f16(*s);
2838 }
2839}
2840
2841pub fn batch_f32_to_u8(src: &[f32], dst: &mut [u8]) {
2843 debug_assert_eq!(src.len(), dst.len());
2844 for (s, d) in src.iter().zip(dst.iter_mut()) {
2845 *d = f32_to_u8_saturating(*s);
2846 }
2847}
2848
2849#[cfg(target_arch = "aarch64")]
2854#[allow(unsafe_op_in_unsafe_fn)]
2855mod neon_quant {
2856 use std::arch::aarch64::*;
2857
2858 #[allow(clippy::incompatible_msrv)]
2864 #[target_feature(enable = "neon")]
2865 pub unsafe fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2866 let chunks16 = dim / 16;
2867 let remainder = dim % 16;
2868
2869 let mut acc_dot0 = vdupq_n_f32(0.0);
2871 let mut acc_dot1 = vdupq_n_f32(0.0);
2872 let mut acc_norm0 = vdupq_n_f32(0.0);
2873 let mut acc_norm1 = vdupq_n_f32(0.0);
2874
2875 for c in 0..chunks16 {
2876 let base = c * 16;
2877
2878 let v_raw0 = vld1q_u16(vec_f16.as_ptr().add(base));
2880 let v_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw0)));
2881 let v_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw0)));
2882 let q_raw0 = vld1q_u16(query_f16.as_ptr().add(base));
2883 let q_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw0)));
2884 let q_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw0)));
2885
2886 acc_dot0 = vfmaq_f32(acc_dot0, q_lo0, v_lo0);
2887 acc_dot0 = vfmaq_f32(acc_dot0, q_hi0, v_hi0);
2888 acc_norm0 = vfmaq_f32(acc_norm0, v_lo0, v_lo0);
2889 acc_norm0 = vfmaq_f32(acc_norm0, v_hi0, v_hi0);
2890
2891 let v_raw1 = vld1q_u16(vec_f16.as_ptr().add(base + 8));
2893 let v_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw1)));
2894 let v_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw1)));
2895 let q_raw1 = vld1q_u16(query_f16.as_ptr().add(base + 8));
2896 let q_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw1)));
2897 let q_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw1)));
2898
2899 acc_dot1 = vfmaq_f32(acc_dot1, q_lo1, v_lo1);
2900 acc_dot1 = vfmaq_f32(acc_dot1, q_hi1, v_hi1);
2901 acc_norm1 = vfmaq_f32(acc_norm1, v_lo1, v_lo1);
2902 acc_norm1 = vfmaq_f32(acc_norm1, v_hi1, v_hi1);
2903 }
2904
2905 let mut dot = vaddvq_f32(vaddq_f32(acc_dot0, acc_dot1));
2907 let mut norm = vaddvq_f32(vaddq_f32(acc_norm0, acc_norm1));
2908
2909 let base = chunks16 * 16;
2911 for i in 0..remainder {
2912 let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
2913 let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
2914 dot += q * v;
2915 norm += v * v;
2916 }
2917
2918 (dot, norm)
2919 }
2920
2921 #[target_feature(enable = "neon")]
2924 pub unsafe fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2925 let scale = vdupq_n_f32(super::U8_INV_SCALE);
2926 let offset = vdupq_n_f32(-1.0);
2927
2928 let chunks16 = dim / 16;
2929 let remainder = dim % 16;
2930
2931 let mut acc_dot = vdupq_n_f32(0.0);
2932 let mut acc_norm = vdupq_n_f32(0.0);
2933
2934 for c in 0..chunks16 {
2935 let base = c * 16;
2936
2937 let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
2939
2940 let lo8 = vget_low_u8(bytes);
2942 let hi8 = vget_high_u8(bytes);
2943 let lo16 = vmovl_u8(lo8);
2944 let hi16 = vmovl_u8(hi8);
2945
2946 let f0 = vaddq_f32(
2947 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
2948 offset,
2949 );
2950 let f1 = vaddq_f32(
2951 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
2952 offset,
2953 );
2954 let f2 = vaddq_f32(
2955 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
2956 offset,
2957 );
2958 let f3 = vaddq_f32(
2959 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
2960 offset,
2961 );
2962
2963 let q0 = vld1q_f32(query.as_ptr().add(base));
2964 let q1 = vld1q_f32(query.as_ptr().add(base + 4));
2965 let q2 = vld1q_f32(query.as_ptr().add(base + 8));
2966 let q3 = vld1q_f32(query.as_ptr().add(base + 12));
2967
2968 acc_dot = vfmaq_f32(acc_dot, q0, f0);
2969 acc_dot = vfmaq_f32(acc_dot, q1, f1);
2970 acc_dot = vfmaq_f32(acc_dot, q2, f2);
2971 acc_dot = vfmaq_f32(acc_dot, q3, f3);
2972
2973 acc_norm = vfmaq_f32(acc_norm, f0, f0);
2974 acc_norm = vfmaq_f32(acc_norm, f1, f1);
2975 acc_norm = vfmaq_f32(acc_norm, f2, f2);
2976 acc_norm = vfmaq_f32(acc_norm, f3, f3);
2977 }
2978
2979 let mut dot = vaddvq_f32(acc_dot);
2980 let mut norm = vaddvq_f32(acc_norm);
2981
2982 let base = chunks16 * 16;
2983 for i in 0..remainder {
2984 let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
2985 dot += *query.get_unchecked(base + i) * v;
2986 norm += v * v;
2987 }
2988
2989 (dot, norm)
2990 }
2991
2992 #[allow(clippy::incompatible_msrv)]
2994 #[target_feature(enable = "neon")]
2995 pub unsafe fn dot_product_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2996 let chunks8 = dim / 8;
2997 let remainder = dim % 8;
2998
2999 let mut acc = vdupq_n_f32(0.0);
3000
3001 for c in 0..chunks8 {
3002 let base = c * 8;
3003 let v_raw = vld1q_u16(vec_f16.as_ptr().add(base));
3004 let v_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw)));
3005 let v_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw)));
3006 let q_raw = vld1q_u16(query_f16.as_ptr().add(base));
3007 let q_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw)));
3008 let q_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw)));
3009 acc = vfmaq_f32(acc, q_lo, v_lo);
3010 acc = vfmaq_f32(acc, q_hi, v_hi);
3011 }
3012
3013 let mut dot = vaddvq_f32(acc);
3014 let base = chunks8 * 8;
3015 for i in 0..remainder {
3016 let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
3017 let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
3018 dot += q * v;
3019 }
3020 dot
3021 }
3022
3023 #[target_feature(enable = "neon")]
3025 pub unsafe fn dot_product_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3026 let scale = vdupq_n_f32(super::U8_INV_SCALE);
3027 let offset = vdupq_n_f32(-1.0);
3028 let chunks16 = dim / 16;
3029 let remainder = dim % 16;
3030
3031 let mut acc = vdupq_n_f32(0.0);
3032
3033 for c in 0..chunks16 {
3034 let base = c * 16;
3035 let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
3036 let lo8 = vget_low_u8(bytes);
3037 let hi8 = vget_high_u8(bytes);
3038 let lo16 = vmovl_u8(lo8);
3039 let hi16 = vmovl_u8(hi8);
3040 let f0 = vaddq_f32(
3041 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
3042 offset,
3043 );
3044 let f1 = vaddq_f32(
3045 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
3046 offset,
3047 );
3048 let f2 = vaddq_f32(
3049 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
3050 offset,
3051 );
3052 let f3 = vaddq_f32(
3053 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
3054 offset,
3055 );
3056 let q0 = vld1q_f32(query.as_ptr().add(base));
3057 let q1 = vld1q_f32(query.as_ptr().add(base + 4));
3058 let q2 = vld1q_f32(query.as_ptr().add(base + 8));
3059 let q3 = vld1q_f32(query.as_ptr().add(base + 12));
3060 acc = vfmaq_f32(acc, q0, f0);
3061 acc = vfmaq_f32(acc, q1, f1);
3062 acc = vfmaq_f32(acc, q2, f2);
3063 acc = vfmaq_f32(acc, q3, f3);
3064 }
3065
3066 let mut dot = vaddvq_f32(acc);
3067 let base = chunks16 * 16;
3068 for i in 0..remainder {
3069 let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
3070 dot += *query.get_unchecked(base + i) * v;
3071 }
3072 dot
3073 }
3074}
3075
3076fn fused_dot_norm_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
3081 (0..dim).fold((0.0f32, 0.0f32), |(dot, norm), i| {
3082 let v = f16_to_f32(vec_f16[i]);
3083 let q = f16_to_f32(query_f16[i]);
3084 (
3085 dot.algebraic_add(q.algebraic_mul(v)),
3086 norm.algebraic_add(v.algebraic_mul(v)),
3087 )
3088 })
3089}
3090
3091fn fused_dot_norm_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
3092 (0..dim).fold((0.0f32, 0.0f32), |(dot, norm), i| {
3093 let v = u8_to_f32(vec_u8[i]);
3094 (
3095 dot.algebraic_add(query[i].algebraic_mul(v)),
3096 norm.algebraic_add(v.algebraic_mul(v)),
3097 )
3098 })
3099}
3100
3101fn dot_product_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
3102 (0..dim).fold(0.0f32, |dot, i| {
3103 dot.algebraic_add(f16_to_f32(query_f16[i]).algebraic_mul(f16_to_f32(vec_f16[i])))
3104 })
3105}
3106
3107fn dot_product_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3108 (0..dim).fold(0.0f32, |dot, i| {
3109 dot.algebraic_add(query[i].algebraic_mul(u8_to_f32(vec_u8[i])))
3110 })
3111}
3112
3113#[cfg(target_arch = "x86_64")]
3118#[target_feature(enable = "sse2", enable = "sse4.1")]
3119#[allow(unsafe_op_in_unsafe_fn)]
3120unsafe fn fused_dot_norm_f16_sse(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
3121 use std::arch::x86_64::*;
3122
3123 let chunks = dim / 4;
3124 let remainder = dim % 4;
3125
3126 let mut acc_dot = _mm_setzero_ps();
3127 let mut acc_norm = _mm_setzero_ps();
3128
3129 for chunk in 0..chunks {
3130 let base = chunk * 4;
3131 let v0 = f16_to_f32(*vec_f16.get_unchecked(base));
3133 let v1 = f16_to_f32(*vec_f16.get_unchecked(base + 1));
3134 let v2 = f16_to_f32(*vec_f16.get_unchecked(base + 2));
3135 let v3 = f16_to_f32(*vec_f16.get_unchecked(base + 3));
3136 let vb = _mm_set_ps(v3, v2, v1, v0);
3137
3138 let q0 = f16_to_f32(*query_f16.get_unchecked(base));
3139 let q1 = f16_to_f32(*query_f16.get_unchecked(base + 1));
3140 let q2 = f16_to_f32(*query_f16.get_unchecked(base + 2));
3141 let q3 = f16_to_f32(*query_f16.get_unchecked(base + 3));
3142 let va = _mm_set_ps(q3, q2, q1, q0);
3143
3144 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
3145 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
3146 }
3147
3148 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
3150 let sums_d = _mm_add_ps(acc_dot, shuf_d);
3151 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
3152 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
3153
3154 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
3155 let sums_n = _mm_add_ps(acc_norm, shuf_n);
3156 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
3157 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
3158
3159 let base = chunks * 4;
3160 for i in 0..remainder {
3161 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
3162 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
3163 dot += q * v;
3164 norm += v * v;
3165 }
3166
3167 (dot, norm)
3168}
3169
3170#[cfg(target_arch = "x86_64")]
3171#[target_feature(enable = "sse2", enable = "sse4.1")]
3172#[allow(unsafe_op_in_unsafe_fn)]
3173unsafe fn fused_dot_norm_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
3174 use std::arch::x86_64::*;
3175
3176 let scale = _mm_set1_ps(U8_INV_SCALE);
3177 let offset = _mm_set1_ps(-1.0);
3178
3179 let chunks = dim / 4;
3180 let remainder = dim % 4;
3181
3182 let mut acc_dot = _mm_setzero_ps();
3183 let mut acc_norm = _mm_setzero_ps();
3184
3185 for chunk in 0..chunks {
3186 let base = chunk * 4;
3187
3188 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
3190 vec_u8.as_ptr().add(base) as *const i32
3191 ));
3192 let ints = _mm_cvtepu8_epi32(bytes);
3193 let floats = _mm_cvtepi32_ps(ints);
3194 let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
3195
3196 let va = _mm_loadu_ps(query.as_ptr().add(base));
3197
3198 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
3199 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
3200 }
3201
3202 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
3204 let sums_d = _mm_add_ps(acc_dot, shuf_d);
3205 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
3206 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
3207
3208 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
3209 let sums_n = _mm_add_ps(acc_norm, shuf_n);
3210 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
3211 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
3212
3213 let base = chunks * 4;
3214 for i in 0..remainder {
3215 let v = u8_to_f32(*vec_u8.get_unchecked(base + i));
3216 dot += *query.get_unchecked(base + i) * v;
3217 norm += v * v;
3218 }
3219
3220 (dot, norm)
3221}
3222
3223#[cfg(target_arch = "x86_64")]
3228#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
3229#[allow(unsafe_op_in_unsafe_fn)]
3230unsafe fn fused_dot_norm_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
3231 use std::arch::x86_64::*;
3232
3233 let chunks16 = dim / 16;
3234 let remainder = dim % 16;
3235
3236 let mut acc_dot0 = _mm256_setzero_ps();
3238 let mut acc_dot1 = _mm256_setzero_ps();
3239 let mut acc_norm0 = _mm256_setzero_ps();
3240 let mut acc_norm1 = _mm256_setzero_ps();
3241
3242 for c in 0..chunks16 {
3243 let base = c * 16;
3244
3245 let v_raw0 = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
3247 let vb0 = _mm256_cvtph_ps(v_raw0);
3248 let q_raw0 = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
3249 let qa0 = _mm256_cvtph_ps(q_raw0);
3250 acc_dot0 = _mm256_fmadd_ps(qa0, vb0, acc_dot0);
3251 acc_norm0 = _mm256_fmadd_ps(vb0, vb0, acc_norm0);
3252
3253 let v_raw1 = _mm_loadu_si128(vec_f16.as_ptr().add(base + 8) as *const __m128i);
3255 let vb1 = _mm256_cvtph_ps(v_raw1);
3256 let q_raw1 = _mm_loadu_si128(query_f16.as_ptr().add(base + 8) as *const __m128i);
3257 let qa1 = _mm256_cvtph_ps(q_raw1);
3258 acc_dot1 = _mm256_fmadd_ps(qa1, vb1, acc_dot1);
3259 acc_norm1 = _mm256_fmadd_ps(vb1, vb1, acc_norm1);
3260 }
3261
3262 let acc_dot = _mm256_add_ps(acc_dot0, acc_dot1);
3264 let acc_norm = _mm256_add_ps(acc_norm0, acc_norm1);
3265
3266 let hi_d = _mm256_extractf128_ps(acc_dot, 1);
3268 let lo_d = _mm256_castps256_ps128(acc_dot);
3269 let sum_d = _mm_add_ps(lo_d, hi_d);
3270 let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
3271 let sums_d = _mm_add_ps(sum_d, shuf_d);
3272 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
3273 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
3274
3275 let hi_n = _mm256_extractf128_ps(acc_norm, 1);
3276 let lo_n = _mm256_castps256_ps128(acc_norm);
3277 let sum_n = _mm_add_ps(lo_n, hi_n);
3278 let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
3279 let sums_n = _mm_add_ps(sum_n, shuf_n);
3280 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
3281 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
3282
3283 let base = chunks16 * 16;
3284 for i in 0..remainder {
3285 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
3286 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
3287 dot += q * v;
3288 norm += v * v;
3289 }
3290
3291 (dot, norm)
3292}
3293
3294#[cfg(target_arch = "x86_64")]
3295#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
3296#[allow(unsafe_op_in_unsafe_fn)]
3297unsafe fn dot_product_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
3298 use std::arch::x86_64::*;
3299
3300 let chunks = dim / 8;
3301 let remainder = dim % 8;
3302 let mut acc = _mm256_setzero_ps();
3303
3304 for chunk in 0..chunks {
3305 let base = chunk * 8;
3306 let v_raw = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
3307 let vb = _mm256_cvtph_ps(v_raw);
3308 let q_raw = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
3309 let qa = _mm256_cvtph_ps(q_raw);
3310 acc = _mm256_fmadd_ps(qa, vb, acc);
3311 }
3312
3313 let hi = _mm256_extractf128_ps(acc, 1);
3314 let lo = _mm256_castps256_ps128(acc);
3315 let sum = _mm_add_ps(lo, hi);
3316 let shuf = _mm_shuffle_ps(sum, sum, 0b10_11_00_01);
3317 let sums = _mm_add_ps(sum, shuf);
3318 let shuf2 = _mm_movehl_ps(sums, sums);
3319 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
3320
3321 let base = chunks * 8;
3322 for i in 0..remainder {
3323 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
3324 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
3325 dot += q * v;
3326 }
3327 dot
3328}
3329
3330#[cfg(target_arch = "x86_64")]
3331#[target_feature(enable = "sse2", enable = "sse4.1")]
3332#[allow(unsafe_op_in_unsafe_fn)]
3333unsafe fn dot_product_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3334 use std::arch::x86_64::*;
3335
3336 let scale = _mm_set1_ps(U8_INV_SCALE);
3337 let offset = _mm_set1_ps(-1.0);
3338 let chunks = dim / 4;
3339 let remainder = dim % 4;
3340 let mut acc = _mm_setzero_ps();
3341
3342 for chunk in 0..chunks {
3343 let base = chunk * 4;
3344 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
3345 vec_u8.as_ptr().add(base) as *const i32
3346 ));
3347 let ints = _mm_cvtepu8_epi32(bytes);
3348 let floats = _mm_cvtepi32_ps(ints);
3349 let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
3350 let va = _mm_loadu_ps(query.as_ptr().add(base));
3351 acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
3352 }
3353
3354 let shuf = _mm_shuffle_ps(acc, acc, 0b10_11_00_01);
3355 let sums = _mm_add_ps(acc, shuf);
3356 let shuf2 = _mm_movehl_ps(sums, sums);
3357 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
3358
3359 let base = chunks * 4;
3360 for i in 0..remainder {
3361 dot += *query.get_unchecked(base + i) * u8_to_f32(*vec_u8.get_unchecked(base + i));
3362 }
3363 dot
3364}
3365
3366#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3372pub enum QuantF16Kernel {
3373 #[cfg(target_arch = "aarch64")]
3374 Neon,
3375 #[cfg(target_arch = "x86_64")]
3376 F16c,
3377 #[cfg(target_arch = "x86_64")]
3380 Sse,
3381 Scalar,
3382}
3383
3384impl QuantF16Kernel {
3385 #[inline]
3386 pub fn resolve() -> Self {
3387 #[cfg(target_arch = "aarch64")]
3388 {
3389 Self::Neon
3390 }
3391 #[cfg(target_arch = "x86_64")]
3392 {
3393 if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("fma") {
3394 return Self::F16c;
3395 }
3396 if sse::is_available() {
3397 return Self::Sse;
3398 }
3399 Self::Scalar
3400 }
3401 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
3402 {
3403 Self::Scalar
3404 }
3405 }
3406
3407 #[inline]
3408 pub fn fused_dot_norm(self, query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
3409 match self {
3410 #[cfg(target_arch = "aarch64")]
3411 Self::Neon => unsafe { neon_quant::fused_dot_norm_f16(query_f16, vec_f16, dim) },
3412 #[cfg(target_arch = "x86_64")]
3413 Self::F16c => unsafe { fused_dot_norm_f16_f16c(query_f16, vec_f16, dim) },
3414 #[cfg(target_arch = "x86_64")]
3415 Self::Sse => unsafe { fused_dot_norm_f16_sse(query_f16, vec_f16, dim) },
3416 Self::Scalar => fused_dot_norm_f16_scalar(query_f16, vec_f16, dim),
3417 }
3418 }
3419
3420 #[inline]
3421 pub fn dot(self, query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
3422 match self {
3423 #[cfg(target_arch = "aarch64")]
3424 Self::Neon => unsafe { neon_quant::dot_product_f16(query_f16, vec_f16, dim) },
3425 #[cfg(target_arch = "x86_64")]
3426 Self::F16c => unsafe { dot_product_f16_f16c(query_f16, vec_f16, dim) },
3427 #[cfg(target_arch = "x86_64")]
3428 Self::Sse => dot_product_f16_scalar(query_f16, vec_f16, dim),
3429 Self::Scalar => dot_product_f16_scalar(query_f16, vec_f16, dim),
3430 }
3431 }
3432}
3433
3434#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3436pub enum QuantU8Kernel {
3437 #[cfg(target_arch = "aarch64")]
3438 Neon,
3439 #[cfg(target_arch = "x86_64")]
3440 Sse,
3441 Scalar,
3442}
3443
3444impl QuantU8Kernel {
3445 #[inline]
3446 pub fn resolve() -> Self {
3447 #[cfg(target_arch = "aarch64")]
3448 {
3449 Self::Neon
3450 }
3451 #[cfg(target_arch = "x86_64")]
3452 {
3453 if sse::is_available() {
3454 return Self::Sse;
3455 }
3456 Self::Scalar
3457 }
3458 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
3459 {
3460 Self::Scalar
3461 }
3462 }
3463
3464 #[inline]
3465 pub fn fused_dot_norm(self, query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
3466 match self {
3467 #[cfg(target_arch = "aarch64")]
3468 Self::Neon => unsafe { neon_quant::fused_dot_norm_u8(query, vec_u8, dim) },
3469 #[cfg(target_arch = "x86_64")]
3470 Self::Sse => unsafe { fused_dot_norm_u8_sse(query, vec_u8, dim) },
3471 Self::Scalar => fused_dot_norm_u8_scalar(query, vec_u8, dim),
3472 }
3473 }
3474
3475 #[inline]
3476 pub fn dot(self, query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3477 match self {
3478 #[cfg(target_arch = "aarch64")]
3479 Self::Neon => unsafe { neon_quant::dot_product_u8(query, vec_u8, dim) },
3480 #[cfg(target_arch = "x86_64")]
3481 Self::Sse => unsafe { dot_product_u8_sse(query, vec_u8, dim) },
3482 Self::Scalar => dot_product_u8_scalar(query, vec_u8, dim),
3483 }
3484 }
3485}
3486
3487#[inline]
3488fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
3489 QuantF16Kernel::resolve().fused_dot_norm(query_f16, vec_f16, dim)
3490}
3491
3492#[inline]
3493fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
3494 QuantU8Kernel::resolve().fused_dot_norm(query, vec_u8, dim)
3495}
3496
3497#[inline]
3500fn dot_product_f16_quant(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
3501 QuantF16Kernel::resolve().dot(query_f16, vec_f16, dim)
3502}
3503
3504#[inline]
3505fn dot_product_u8_quant(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3506 QuantU8Kernel::resolve().dot(query, vec_u8, dim)
3507}
3508
3509#[inline]
3520pub fn batch_cosine_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3521 let n = scores.len();
3522 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3523 let required = n
3524 .checked_mul(vec_bytes)
3525 .expect("f16 batch byte length overflow");
3526 assert_eq!(
3527 query.len(),
3528 dim,
3529 "f16 batch cosine query dimension mismatch"
3530 );
3531 assert!(
3532 vectors_raw.len() >= required,
3533 "f16 batch cosine vectors are truncated: need {required} bytes, got {}",
3534 vectors_raw.len()
3535 );
3536 if required > 0 {
3537 assert!(
3538 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3539 "f16 batch cosine vectors are not 2-byte aligned"
3540 );
3541 }
3542 if dim == 0 || n == 0 {
3543 return;
3544 }
3545
3546 let norm_q_sq = dot_product_f32(query, query, dim);
3548 if norm_q_sq < f32::EPSILON {
3549 for s in scores.iter_mut() {
3550 *s = 0.0;
3551 }
3552 return;
3553 }
3554 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3555
3556 let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3558
3559 for i in 0..n {
3560 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3561 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3562
3563 let (dot, norm_v_sq) = fused_dot_norm_f16(&query_f16, f16_slice, dim);
3564 scores[i] = if norm_v_sq < f32::EPSILON {
3565 0.0
3566 } else {
3567 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3568 };
3569 }
3570}
3571
3572#[inline]
3579pub fn batch_cosine_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3580 let n = scores.len();
3581 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3582 assert_eq!(query.len(), dim, "u8 batch cosine query dimension mismatch");
3583 assert!(
3584 vectors_raw.len() >= required,
3585 "u8 batch cosine vectors are truncated: need {required} bytes, got {}",
3586 vectors_raw.len()
3587 );
3588 if dim == 0 || n == 0 {
3589 return;
3590 }
3591
3592 let norm_q_sq = dot_product_f32(query, query, dim);
3593 if norm_q_sq < f32::EPSILON {
3594 for s in scores.iter_mut() {
3595 *s = 0.0;
3596 }
3597 return;
3598 }
3599 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3600
3601 for i in 0..n {
3602 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3603
3604 let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3605 scores[i] = if norm_v_sq < f32::EPSILON {
3606 0.0
3607 } else {
3608 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3609 };
3610 }
3611}
3612
3613#[inline]
3622pub fn batch_dot_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
3623 let n = scores.len();
3624 let required = n
3625 .checked_mul(dim)
3626 .expect("batch dot vector length overflow");
3627 assert_eq!(query.len(), dim, "batch dot query dimension mismatch");
3628 assert!(
3629 vectors.len() >= required,
3630 "batch dot vectors are truncated: need {required}, got {}",
3631 vectors.len()
3632 );
3633
3634 if dim == 0 || n == 0 {
3635 return;
3636 }
3637
3638 let norm_q_sq = dot_product_f32(query, query, dim);
3639 if norm_q_sq < f32::EPSILON {
3640 for s in scores.iter_mut() {
3641 *s = 0.0;
3642 }
3643 return;
3644 }
3645 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3646
3647 for i in 0..n {
3648 let vec = &vectors[i * dim..(i + 1) * dim];
3649 let dot = dot_product_f32(query, vec, dim);
3650 scores[i] = dot * inv_norm_q;
3651 }
3652}
3653
3654#[inline]
3659pub fn batch_dot_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3660 let n = scores.len();
3661 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3662 let required = n
3663 .checked_mul(vec_bytes)
3664 .expect("f16 batch byte length overflow");
3665 assert_eq!(query.len(), dim, "f16 batch dot query dimension mismatch");
3666 assert!(
3667 vectors_raw.len() >= required,
3668 "f16 batch dot vectors are truncated: need {required} bytes, got {}",
3669 vectors_raw.len()
3670 );
3671 if required > 0 {
3672 assert!(
3673 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3674 "f16 batch dot vectors are not 2-byte aligned"
3675 );
3676 }
3677 if dim == 0 || n == 0 {
3678 return;
3679 }
3680
3681 let norm_q_sq = dot_product_f32(query, query, dim);
3682 if norm_q_sq < f32::EPSILON {
3683 for s in scores.iter_mut() {
3684 *s = 0.0;
3685 }
3686 return;
3687 }
3688 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3689
3690 let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3691 for i in 0..n {
3692 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3693 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3694 let dot = dot_product_f16_quant(&query_f16, f16_slice, dim);
3695 scores[i] = dot * inv_norm_q;
3696 }
3697}
3698
3699#[inline]
3704pub fn batch_dot_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3705 let n = scores.len();
3706 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3707 assert_eq!(query.len(), dim, "u8 batch dot query dimension mismatch");
3708 assert!(
3709 vectors_raw.len() >= required,
3710 "u8 batch dot vectors are truncated: need {required} bytes, got {}",
3711 vectors_raw.len()
3712 );
3713 if dim == 0 || n == 0 {
3714 return;
3715 }
3716
3717 let norm_q_sq = dot_product_f32(query, query, dim);
3718 if norm_q_sq < f32::EPSILON {
3719 for s in scores.iter_mut() {
3720 *s = 0.0;
3721 }
3722 return;
3723 }
3724 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3725
3726 for i in 0..n {
3727 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3728 let dot = dot_product_u8_quant(query, u8_slice, dim);
3729 scores[i] = dot * inv_norm_q;
3730 }
3731}
3732
3733#[inline]
3739pub fn batch_cosine_scores_precomp(
3740 query: &[f32],
3741 vectors: &[f32],
3742 dim: usize,
3743 scores: &mut [f32],
3744 inv_norm_q: f32,
3745) {
3746 let n = scores.len();
3747 let required = n
3748 .checked_mul(dim)
3749 .expect("precomputed cosine vector length overflow");
3750 assert_eq!(
3751 query.len(),
3752 dim,
3753 "precomputed cosine query dimension mismatch"
3754 );
3755 assert!(
3756 vectors.len() >= required,
3757 "precomputed cosine vectors are truncated: need {required}, got {}",
3758 vectors.len()
3759 );
3760 #[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
3765 macro_rules! score_simd_rows {
3766 ($kernel:path) => {{
3767 for i in 0..n {
3768 let vec = &vectors[i * dim..(i + 1) * dim];
3769 let (dot, norm_v_sq) = unsafe { $kernel(query, vec, dim) };
3770 scores[i] = if norm_v_sq < f32::EPSILON {
3771 0.0
3772 } else {
3773 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3774 };
3775 }
3776 }};
3777 }
3778 match DenseF32Kernel::resolve() {
3779 #[cfg(target_arch = "aarch64")]
3780 DenseF32Kernel::Neon => score_simd_rows!(fused_dot_norm_neon),
3781 #[cfg(target_arch = "x86_64")]
3782 DenseF32Kernel::Avx512 => score_simd_rows!(fused_dot_norm_avx512),
3783 #[cfg(target_arch = "x86_64")]
3784 DenseF32Kernel::Avx2Fma => score_simd_rows!(fused_dot_norm_avx2),
3785 #[cfg(target_arch = "x86_64")]
3786 DenseF32Kernel::Sse => score_simd_rows!(fused_dot_norm_sse),
3787 DenseF32Kernel::Scalar => {
3788 for i in 0..n {
3789 let vec = &vectors[i * dim..(i + 1) * dim];
3790 let (dot, norm_v_sq) = fused_dot_norm_scalar(query, vec);
3791 scores[i] = if norm_v_sq < f32::EPSILON {
3792 0.0
3793 } else {
3794 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3795 };
3796 }
3797 }
3798 }
3799}
3800
3801#[inline]
3803pub fn batch_cosine_scores_f16_precomp(
3804 query_f16: &[u16],
3805 vectors_raw: &[u8],
3806 dim: usize,
3807 scores: &mut [f32],
3808 inv_norm_q: f32,
3809) {
3810 let n = scores.len();
3811 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3812 let required = n
3813 .checked_mul(vec_bytes)
3814 .expect("precomputed f16 cosine batch byte length overflow");
3815 assert_eq!(
3816 query_f16.len(),
3817 dim,
3818 "precomputed f16 cosine query dimension mismatch"
3819 );
3820 assert!(
3821 vectors_raw.len() >= required,
3822 "precomputed f16 cosine vectors are truncated: need {required} bytes, got {}",
3823 vectors_raw.len()
3824 );
3825 if required > 0 {
3826 assert!(
3827 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3828 "precomputed f16 cosine vectors are not 2-byte aligned"
3829 );
3830 }
3831 let kernel = QuantF16Kernel::resolve();
3832 for i in 0..n {
3833 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3834 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3835 let (dot, norm_v_sq) = kernel.fused_dot_norm(query_f16, f16_slice, dim);
3836 scores[i] = if norm_v_sq < f32::EPSILON {
3837 0.0
3838 } else {
3839 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3840 };
3841 }
3842}
3843
3844#[inline]
3846pub fn batch_cosine_scores_u8_precomp(
3847 query: &[f32],
3848 vectors_raw: &[u8],
3849 dim: usize,
3850 scores: &mut [f32],
3851 inv_norm_q: f32,
3852) {
3853 let n = scores.len();
3854 let required = n
3855 .checked_mul(dim)
3856 .expect("precomputed u8 cosine batch byte length overflow");
3857 assert_eq!(
3858 query.len(),
3859 dim,
3860 "precomputed u8 cosine query dimension mismatch"
3861 );
3862 assert!(
3863 vectors_raw.len() >= required,
3864 "precomputed u8 cosine vectors are truncated: need {required} bytes, got {}",
3865 vectors_raw.len()
3866 );
3867 let kernel = QuantU8Kernel::resolve();
3868 for i in 0..n {
3869 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3870 let (dot, norm_v_sq) = kernel.fused_dot_norm(query, u8_slice, dim);
3871 scores[i] = if norm_v_sq < f32::EPSILON {
3872 0.0
3873 } else {
3874 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3875 };
3876 }
3877}
3878
3879#[inline]
3881pub fn batch_dot_scores_precomp(
3882 query: &[f32],
3883 vectors: &[f32],
3884 dim: usize,
3885 scores: &mut [f32],
3886 inv_norm_q: f32,
3887) {
3888 let n = scores.len();
3889 let required = n
3890 .checked_mul(dim)
3891 .expect("precomputed dot vector length overflow");
3892 assert_eq!(query.len(), dim, "precomputed dot query dimension mismatch");
3893 assert!(
3894 vectors.len() >= required,
3895 "precomputed dot vectors are truncated: need {required}, got {}",
3896 vectors.len()
3897 );
3898 let kernel = DenseF32Kernel::resolve();
3899 for i in 0..n {
3900 let vec = &vectors[i * dim..(i + 1) * dim];
3901 scores[i] = kernel.dot(query, vec, dim) * inv_norm_q;
3902 }
3903}
3904
3905#[inline]
3907pub fn batch_dot_scores_f16_precomp(
3908 query_f16: &[u16],
3909 vectors_raw: &[u8],
3910 dim: usize,
3911 scores: &mut [f32],
3912 inv_norm_q: f32,
3913) {
3914 let n = scores.len();
3915 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3916 let required = n
3917 .checked_mul(vec_bytes)
3918 .expect("precomputed f16 dot batch byte length overflow");
3919 assert_eq!(
3920 query_f16.len(),
3921 dim,
3922 "precomputed f16 dot query dimension mismatch"
3923 );
3924 assert!(
3925 vectors_raw.len() >= required,
3926 "precomputed f16 dot vectors are truncated: need {required} bytes, got {}",
3927 vectors_raw.len()
3928 );
3929 if required > 0 {
3930 assert!(
3931 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3932 "precomputed f16 dot vectors are not 2-byte aligned"
3933 );
3934 }
3935 let kernel = QuantF16Kernel::resolve();
3936 for i in 0..n {
3937 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3938 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3939 scores[i] = kernel.dot(query_f16, f16_slice, dim) * inv_norm_q;
3940 }
3941}
3942
3943#[inline]
3945pub fn batch_dot_scores_u8_precomp(
3946 query: &[f32],
3947 vectors_raw: &[u8],
3948 dim: usize,
3949 scores: &mut [f32],
3950 inv_norm_q: f32,
3951) {
3952 let n = scores.len();
3953 let required = n
3954 .checked_mul(dim)
3955 .expect("precomputed u8 dot batch byte length overflow");
3956 assert_eq!(
3957 query.len(),
3958 dim,
3959 "precomputed u8 dot query dimension mismatch"
3960 );
3961 assert!(
3962 vectors_raw.len() >= required,
3963 "precomputed u8 dot vectors are truncated: need {required} bytes, got {}",
3964 vectors_raw.len()
3965 );
3966 let kernel = QuantU8Kernel::resolve();
3967 for i in 0..n {
3968 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3969 scores[i] = kernel.dot(query, u8_slice, dim) * inv_norm_q;
3970 }
3971}
3972
3973#[inline]
3978pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
3979 assert_eq!(a.len(), b.len(), "cosine vector dimension mismatch");
3980 let count = a.len();
3981
3982 if count == 0 {
3983 return 0.0;
3984 }
3985
3986 let dot = dot_product_f32(a, b, count);
3987 let norm_a = dot_product_f32(a, a, count);
3988 let norm_b = dot_product_f32(b, b, count);
3989
3990 let denom = (norm_a * norm_b).sqrt();
3991 if denom < f32::EPSILON {
3992 return 0.0;
3993 }
3994
3995 dot / denom
3996}
3997
3998#[cfg(target_arch = "x86_64")]
4007#[target_feature(enable = "avx512f,avx512vpopcntdq")]
4008#[allow(unsafe_op_in_unsafe_fn)]
4009unsafe fn hamming_distance_avx512(a: &[u8], b: &[u8]) -> u32 {
4010 use std::arch::x86_64::*;
4011
4012 let len = a.len();
4013 let chunks64 = len / 64;
4014 let mut acc = _mm512_setzero_si512();
4015
4016 for c in 0..chunks64 {
4017 let off = c * 64;
4018 let va = _mm512_loadu_si512(a.as_ptr().add(off) as *const __m512i);
4019 let vb = _mm512_loadu_si512(b.as_ptr().add(off) as *const __m512i);
4020 acc = _mm512_add_epi64(acc, _mm512_popcnt_epi64(_mm512_xor_si512(va, vb)));
4021 }
4022
4023 let base = chunks64 * 64;
4024 _mm512_reduce_add_epi64(acc) as u32 + hamming_distance_scalar(&a[base..], &b[base..])
4025}
4026
4027#[cfg(target_arch = "x86_64")]
4029#[target_feature(enable = "avx512f,avx512vpopcntdq")]
4030#[allow(unsafe_op_in_unsafe_fn)]
4031#[inline]
4032unsafe fn hamming_distance_x4_avx512(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
4033 use std::arch::x86_64::*;
4034
4035 let len = query.len();
4036 let chunks64 = len / 64;
4037 let mut acc = [_mm512_setzero_si512(); 4];
4038
4039 for c in 0..chunks64 {
4040 let off = c * 64;
4041 let vq = _mm512_loadu_si512(query.as_ptr().add(off) as *const __m512i);
4042 for r in 0..4 {
4043 let vr = _mm512_loadu_si512(rows[r].as_ptr().add(off) as *const __m512i);
4044 acc[r] = _mm512_add_epi64(acc[r], _mm512_popcnt_epi64(_mm512_xor_si512(vq, vr)));
4045 }
4046 }
4047
4048 let base = chunks64 * 64;
4049 let tail = &query[base..];
4050 [
4051 _mm512_reduce_add_epi64(acc[0]) as u32 + hamming_distance_scalar(tail, &rows[0][base..]),
4052 _mm512_reduce_add_epi64(acc[1]) as u32 + hamming_distance_scalar(tail, &rows[1][base..]),
4053 _mm512_reduce_add_epi64(acc[2]) as u32 + hamming_distance_scalar(tail, &rows[2][base..]),
4054 _mm512_reduce_add_epi64(acc[3]) as u32 + hamming_distance_scalar(tail, &rows[3][base..]),
4055 ]
4056}
4057
4058#[inline]
4060fn hamming_distance_x4_scalar(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
4061 let len = query.len();
4062 let chunks = len / 8;
4063 let mut total = [0u32; 4];
4064
4065 for i in 0..chunks {
4066 let off = i * 8;
4067 let vq = unsafe { std::ptr::read_unaligned(query.as_ptr().add(off) as *const u64) };
4068 for r in 0..4 {
4069 let vr = unsafe { std::ptr::read_unaligned(rows[r].as_ptr().add(off) as *const u64) };
4070 total[r] += (vq ^ vr).count_ones();
4071 }
4072 }
4073
4074 let base = chunks * 8;
4075 for k in base..len {
4076 let q = query[k];
4077 for r in 0..4 {
4078 total[r] += (q ^ rows[r][k]).count_ones();
4079 }
4080 }
4081
4082 total
4083}
4084
4085const HAMMING_ROWS_PER_KERNEL: usize = 4;
4089
4090#[derive(Clone, Copy, Debug, PartialEq, Eq)]
4097pub enum HammingKernel {
4098 #[cfg(target_arch = "x86_64")]
4099 Avx512,
4100 #[cfg(target_arch = "x86_64")]
4101 Avx2,
4102 #[cfg(target_arch = "aarch64")]
4103 Neon,
4104 Scalar,
4105}
4106
4107impl HammingKernel {
4108 #[inline]
4110 pub fn resolve() -> Self {
4111 #[cfg(target_arch = "x86_64")]
4112 {
4113 if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512vpopcntdq") {
4114 return Self::Avx512;
4115 }
4116 if avx2::is_available() {
4117 return Self::Avx2;
4118 }
4119 Self::Scalar
4120 }
4121
4122 #[cfg(target_arch = "aarch64")]
4123 {
4124 Self::Neon
4125 }
4126
4127 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
4128 {
4129 Self::Scalar
4130 }
4131 }
4132
4133 #[inline]
4136 fn vector_bytes(self) -> usize {
4137 match self {
4138 #[cfg(target_arch = "x86_64")]
4139 Self::Avx512 => 64,
4140 #[cfg(target_arch = "x86_64")]
4141 Self::Avx2 => 32,
4142 #[cfg(target_arch = "aarch64")]
4143 Self::Neon => 16,
4144 Self::Scalar => 8,
4145 }
4146 }
4147
4148 #[inline]
4158 fn for_byte_len(self, byte_len: usize) -> Self {
4159 if byte_len < self.vector_bytes() {
4160 Self::Scalar
4161 } else {
4162 self
4163 }
4164 }
4165
4166 #[inline]
4168 pub fn distance(self, a: &[u8], b: &[u8]) -> u32 {
4169 debug_assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
4170 match self.for_byte_len(a.len()) {
4171 #[cfg(target_arch = "x86_64")]
4172 Self::Avx512 => unsafe { hamming_distance_avx512(a, b) },
4173 #[cfg(target_arch = "x86_64")]
4174 Self::Avx2 => unsafe { avx2::hamming_distance(a, b) },
4175 #[cfg(target_arch = "aarch64")]
4176 Self::Neon => unsafe { neon::hamming_distance(a, b) },
4177 Self::Scalar => hamming_distance_scalar(a, b),
4178 }
4179 }
4180
4181 pub fn distances(self, query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
4183 if byte_len == 32 {
4186 self.score_rows(query, db, 32, out, |index| index);
4187 } else {
4188 self.score_rows(query, db, byte_len, out, |index| index);
4189 }
4190 }
4191
4192 pub fn gather_distances(
4197 self,
4198 query: &[u8],
4199 db: &[u8],
4200 byte_len: usize,
4201 ids: &[u32],
4202 out: &mut [u32],
4203 ) {
4204 assert_eq!(
4205 ids.len(),
4206 out.len(),
4207 "Hamming gather needs one output slot per row id"
4208 );
4209 self.score_rows(query, db, byte_len, out, |index| ids[index] as usize);
4210 }
4211
4212 #[inline]
4213 fn score_rows(
4214 self,
4215 query: &[u8],
4216 db: &[u8],
4217 byte_len: usize,
4218 out: &mut [u32],
4219 index_of: impl Fn(usize) -> usize,
4220 ) {
4221 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
4222 if byte_len == 0 || out.is_empty() {
4223 return;
4224 }
4225 let row = |index: usize| -> &[u8] {
4226 let start = index * byte_len;
4227 &db[start..start + byte_len]
4228 };
4229 let kernel = self.for_byte_len(byte_len);
4230 macro_rules! score_with {
4231 ($one:expr, $four:expr) => {{
4232 let mut i = 0;
4233 while i + HAMMING_ROWS_PER_KERNEL <= out.len() {
4234 let quad = [
4235 row(index_of(i)),
4236 row(index_of(i + 1)),
4237 row(index_of(i + 2)),
4238 row(index_of(i + 3)),
4239 ];
4240 out[i..i + HAMMING_ROWS_PER_KERNEL].copy_from_slice(&$four(query, quad));
4241 i += HAMMING_ROWS_PER_KERNEL;
4242 }
4243 while i < out.len() {
4244 out[i] = $one(query, row(index_of(i)));
4245 i += 1;
4246 }
4247 }};
4248 }
4249 match kernel {
4250 #[cfg(target_arch = "x86_64")]
4251 Self::Avx512 => score_with!(
4252 |query, row| unsafe { hamming_distance_avx512(query, row) },
4253 |query, rows| unsafe { hamming_distance_x4_avx512(query, rows) }
4254 ),
4255 #[cfg(target_arch = "x86_64")]
4256 Self::Avx2 => score_with!(
4257 |query, row| unsafe { avx2::hamming_distance(query, row) },
4258 |query, rows| unsafe { avx2::hamming_distance_x4(query, rows) }
4259 ),
4260 #[cfg(target_arch = "aarch64")]
4261 Self::Neon => score_with!(
4262 |query, row| unsafe { neon::hamming_distance(query, row) },
4263 |query, rows| unsafe { neon::hamming_distance_x4(query, rows) }
4264 ),
4265 Self::Scalar => score_with!(hamming_distance_scalar, hamming_distance_x4_scalar),
4266 }
4267 }
4268}
4269
4270#[inline]
4277pub fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
4278 assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
4279 HammingKernel::resolve().distance(a, b)
4280}
4281
4282#[inline]
4285fn hamming_distance_scalar(a: &[u8], b: &[u8]) -> u32 {
4286 let len = a.len();
4287 let chunks = len / 8;
4288 let remainder = len % 8;
4289 let mut total = 0u32;
4290
4291 for i in 0..chunks {
4292 let off = i * 8;
4293 let va = unsafe { std::ptr::read_unaligned(a.as_ptr().add(off) as *const u64) };
4294 let vb = unsafe { std::ptr::read_unaligned(b.as_ptr().add(off) as *const u64) };
4295 total += (va ^ vb).count_ones();
4296 }
4297
4298 let base = chunks * 8;
4299 for i in 0..remainder {
4300 total += (a[base + i] ^ b[base + i]).count_ones();
4301 }
4302
4303 total
4304}
4305
4306pub fn batch_hamming_scores(
4312 query: &[u8],
4313 db: &[u8],
4314 byte_len: usize,
4315 dim_bits: usize,
4316 scores: &mut [f32],
4317) {
4318 let n = scores.len();
4319 let required = n
4320 .checked_mul(byte_len)
4321 .expect("Hamming batch byte length overflow");
4322 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
4323 assert!(
4324 db.len() >= required,
4325 "Hamming batch is truncated: need {required} bytes, got {}",
4326 db.len()
4327 );
4328
4329 if byte_len == 0 || n == 0 || dim_bits == 0 {
4330 return;
4331 }
4332
4333 scores_from_hamming(
4334 HammingKernel::resolve(),
4335 query,
4336 db,
4337 byte_len,
4338 dim_bits,
4339 scores,
4340 );
4341}
4342
4343pub fn scores_from_hamming(
4348 kernel: HammingKernel,
4349 query: &[u8],
4350 db: &[u8],
4351 byte_len: usize,
4352 dim_bits: usize,
4353 scores: &mut [f32],
4354) {
4355 if byte_len == 0 || scores.is_empty() || dim_bits == 0 {
4356 return;
4357 }
4358 let inv_dim = 1.0 / dim_bits as f32;
4359 let mut distances = [0u32; HAMMING_DISTANCE_BLOCK];
4362 for (block_index, block) in scores.chunks_mut(HAMMING_DISTANCE_BLOCK).enumerate() {
4363 let rows = &mut distances[..block.len()];
4364 kernel.distances(
4365 query,
4366 &db[block_index * HAMMING_DISTANCE_BLOCK * byte_len..],
4367 byte_len,
4368 rows,
4369 );
4370 for (score, &distance) in block.iter_mut().zip(rows.iter()) {
4371 *score = 1.0 - distance as f32 * inv_dim;
4372 }
4373 }
4374}
4375
4376const HAMMING_DISTANCE_BLOCK: usize = 64;
4378
4379pub fn batch_hamming_distances(query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
4384 HammingKernel::resolve().distances(query, db, byte_len, out);
4385}
4386
4387#[cfg(test)]
4388mod tests {
4389 #[test]
4390 fn fixed_block_seek_preserves_suffix_lower_bounds_at_unsigned_extremes() {
4391 for length in 0..=128 {
4392 for base in [0u32, 1 << 31, u32::MAX - 512] {
4393 let docs: Vec<_> = (0..length).map(|i| base + i as u32 * 3).collect();
4394 let targets = [0, base, base + 1, base + 127, base + 383, u32::MAX];
4395 for from in 0..=length {
4396 for target in targets {
4397 assert_eq!(
4398 super::find_first_ge_block_from(&docs, from, target),
4399 from + docs[from..].partition_point(|&doc| doc < target),
4400 "length={length} from={from} target={target}"
4401 );
4402 }
4403 }
4404 }
4405 }
4406 for from in 0..=128 {
4407 for target in [0, 7, 8, u32::MAX] {
4408 let docs = [7; 128];
4409 assert_eq!(
4410 super::find_first_ge_block_from(&docs, from, target),
4411 from + docs[from..].partition_point(|&doc| doc < target)
4412 );
4413 }
4414 }
4415 }
4416
4417 #[test]
4418 fn posting_block_intersection_preserves_suffixes_partial_outputs_and_unsigned_ids() {
4419 for base in [0u32, 1 << 31, u32::MAX - 4096] {
4420 for trial in 0..32 {
4421 let all_left: Vec<_> = (0..128).map(|i| base + i * (trial % 7 + 1)).collect();
4422 let all_right: Vec<_> = (0..128)
4423 .map(|i| base + i * (trial % 11 + 1) + trial % 3)
4424 .collect();
4425 for len_a in [0, 1, 7, 8, 9, 127, 128] {
4426 for len_b in [0, 1, 7, 8, 9, 127, 128] {
4427 let left = &all_left[..len_a];
4428 let right = &all_right[..len_b];
4429 for (from_a, from_b) in [(0, 0), (len_a / 2, len_b / 3), (len_a, len_b)] {
4430 let expected: Vec<_> = left[from_a..]
4431 .iter()
4432 .copied()
4433 .filter(|value| right[from_b..].binary_search(value).is_ok())
4434 .collect();
4435 for limit in [1, 7, 128] {
4436 let (mut a, mut b) = (from_a, from_b);
4437 let mut actual = Vec::new();
4438 let mut pairs = [(0u8, 0u8); 128];
4439 while a < len_a && b < len_b {
4440 let previous = (a, b);
4441 let count = super::intersect_posting_blocks(
4442 left,
4443 &mut a,
4444 right,
4445 &mut b,
4446 &mut pairs[..limit],
4447 );
4448 assert!(a > previous.0 || b > previous.1);
4449 for &(l, r) in &pairs[..count] {
4450 assert_eq!(left[l as usize], right[r as usize]);
4451 actual.push(left[l as usize]);
4452 }
4453 }
4454 assert_eq!(actual, expected);
4455 }
4456 }
4457 }
4458 }
4459 }
4460 }
4461 }
4462
4463 #[test]
4467 fn scalar_delta_decode_with_offset_matches_naive_reference_for_all_counts() {
4468 fn naive<const OFFSET: u32>(deltas: &[u32], first: u32) -> Vec<u32> {
4469 let mut out = vec![first];
4470 for &d in deltas {
4471 out.push(out.last().unwrap().wrapping_add(d).wrapping_add(OFFSET));
4472 }
4473 out
4474 }
4475 for count in 0..=257usize {
4476 let deltas: Vec<u32> = (0..count.saturating_sub(1))
4477 .map(|i| [0, 1, 255, 65535, 42, 17][i % 6])
4478 .collect();
4479 for first in [0u32, 7, u32::MAX - 3] {
4480 for bytes in [1usize, 2] {
4481 let mask = if bytes == 1 { 0xFF } else { 0xFFFF };
4482 let masked: Vec<u32> = deltas.iter().map(|d| d & mask).collect();
4483 let mut input = Vec::new();
4484 for d in &masked {
4485 input.extend_from_slice(&d.to_le_bytes()[..bytes]);
4486 }
4487 let (expected0, expected1) = if count == 0 {
4488 (Vec::new(), Vec::new())
4489 } else {
4490 (naive::<0>(&masked, first), naive::<1>(&masked, first))
4491 };
4492 let mut out0 = vec![0xDEAD_BEEF; count + 2];
4493 let mut out1 = vec![0xDEAD_BEEF; count + 2];
4494 if bytes == 1 {
4495 super::scalar::delta_decode_with_offset::<0, 1>(
4496 &input,
4497 &mut out0[..count],
4498 first,
4499 count,
4500 );
4501 super::scalar::delta_decode_with_offset::<1, 1>(
4502 &input,
4503 &mut out1[..count],
4504 first,
4505 count,
4506 );
4507 } else {
4508 super::scalar::delta_decode_with_offset::<0, 2>(
4509 &input,
4510 &mut out0[..count],
4511 first,
4512 count,
4513 );
4514 super::scalar::delta_decode_with_offset::<1, 2>(
4515 &input,
4516 &mut out1[..count],
4517 first,
4518 count,
4519 );
4520 }
4521 assert_eq!(
4522 &out0[..count],
4523 expected0,
4524 "offset 0 bytes={bytes} count={count}"
4525 );
4526 assert_eq!(
4527 &out1[..count],
4528 expected1,
4529 "offset 1 bytes={bytes} count={count}"
4530 );
4531 assert_eq!(&out0[count..], &[0xDEAD_BEEF; 2]);
4532 assert_eq!(&out1[count..], &[0xDEAD_BEEF; 2]);
4533 if bytes == 1 {
4536 let mut simd_out = vec![0; count];
4537 super::unpack_8bit_delta_decode_with_offset::<0>(
4538 &input,
4539 &mut simd_out,
4540 first,
4541 count,
4542 );
4543 assert_eq!(simd_out, expected0, "dispatch 8-bit count={count}");
4544 } else {
4545 let mut simd_out = vec![0; count];
4546 super::unpack_16bit_delta_decode_with_offset::<1>(
4547 &input,
4548 &mut simd_out,
4549 first,
4550 count,
4551 );
4552 assert_eq!(simd_out, expected1, "dispatch 16-bit count={count}");
4553 }
4554 }
4555 }
4556 }
4557 }
4558
4559 #[test]
4560 #[should_panic(expected = "fused delta decode: input holds")]
4561 fn fused_delta_decode_rejects_short_input_before_touching_the_kernels() {
4562 let input = [1u8; 3];
4563 let mut output = [0u32; 8];
4564 super::unpack_8bit_delta_decode(&input, &mut output, 0, 8);
4565 }
4566
4567 #[test]
4568 #[should_panic(expected = "fused delta decode: output holds")]
4569 fn fused_delta_decode_rejects_short_output_before_touching_the_kernels() {
4570 let input = [1u8; 16];
4571 let mut output = [0u32; 4];
4572 super::unpack_16bit_delta_decode(&input, &mut output, 0, 8);
4573 }
4574
4575 #[test]
4576 fn rounded_bit_width_try_from_u8_rejects_unrounded_widths() {
4577 use super::RoundedBitWidth;
4578 assert_eq!(RoundedBitWidth::try_from_u8(0), Some(RoundedBitWidth::Zero));
4579 assert_eq!(
4580 RoundedBitWidth::try_from_u8(8),
4581 Some(RoundedBitWidth::Bits8)
4582 );
4583 assert_eq!(
4584 RoundedBitWidth::try_from_u8(16),
4585 Some(RoundedBitWidth::Bits16)
4586 );
4587 assert_eq!(
4588 RoundedBitWidth::try_from_u8(32),
4589 Some(RoundedBitWidth::Bits32)
4590 );
4591 for bad in (1..=255u8).filter(|b| ![8, 16, 32].contains(b)) {
4592 assert_eq!(RoundedBitWidth::try_from_u8(bad), None, "width {bad}");
4593 }
4594 }
4595
4596 #[test]
4597 fn raw_rounded_gaps_and_legacy_gaps_agree_with_scalar_at_all_tails() {
4598 use super::*;
4599 for (width, mask) in [
4600 (RoundedBitWidth::Zero, 0u32),
4601 (RoundedBitWidth::Bits8, 255),
4602 (RoundedBitWidth::Bits16, 65535),
4603 (RoundedBitWidth::Bits32, u32::MAX),
4604 ] {
4605 for count in 0..=257usize {
4606 let first = u32::MAX - 17;
4607 let gaps: Vec<_> = (0..count.saturating_sub(1))
4608 .map(|i| [0, 1, mask, mask / 2][i % 4] & mask)
4609 .collect();
4610 let mut input = Vec::new();
4611 for &gap in &gaps {
4612 let bytes = gap.to_le_bytes();
4613 input.extend_from_slice(&bytes[..width.bytes_per_value()]);
4614 }
4615 let mut expected = Vec::with_capacity(count);
4616 if count > 0 {
4617 expected.push(first);
4618 for &gap in &gaps {
4619 expected.push(expected.last().unwrap().wrapping_add(gap));
4620 }
4621 }
4622 let mut actual = vec![0xDEADBEEF; count + 4];
4623 unpack_rounded_raw_delta_decode(&input, width, &mut actual[..count], first, count);
4624 assert_eq!(&actual[..count], expected, "width={width:?} count={count}");
4625 assert_eq!(&actual[count..], &[0xDEADBEEF; 4]);
4626 let mut legacy = vec![0; count];
4627 unpack_rounded_delta_decode(&input, width, &mut legacy, first, count);
4628 let biased: Vec<_> = expected
4629 .iter()
4630 .enumerate()
4631 .map(|(i, &doc)| doc.wrapping_add(i as u32))
4632 .collect();
4633 assert_eq!(legacy, biased, "legacy width={width:?} count={count}");
4634 if count == 0 {
4635 continue;
4636 }
4637 #[cfg(target_arch = "x86_64")]
4639 for (available, kernels) in [
4640 (
4641 sse::is_available(),
4642 [
4643 sse::unpack_8bit_delta_decode_with_offset::<0>
4644 as unsafe fn(&[u8], &mut [u32], u32, usize),
4645 sse::unpack_16bit_delta_decode_with_offset::<0>,
4646 ],
4647 ),
4648 (
4649 avx2::is_available(),
4650 [
4651 avx2::unpack_8bit_delta_decode_with_offset::<0>
4652 as unsafe fn(&[u8], &mut [u32], u32, usize),
4653 avx2::unpack_16bit_delta_decode_with_offset::<0>,
4654 ],
4655 ),
4656 ] {
4657 let index = match width {
4658 RoundedBitWidth::Bits8 => Some(0),
4659 RoundedBitWidth::Bits16 => Some(1),
4660 _ => None,
4661 };
4662 if available && let Some(index) = index {
4663 let mut decoded = vec![0; count];
4664 unsafe {
4665 kernels[index](&input, &mut decoded, first, count);
4666 }
4667 assert_eq!(decoded, expected);
4668 }
4669 }
4670 }
4671 }
4672 }
4673 use super::*;
4674
4675 #[test]
4676 fn vector_simd_boundaries_reject_dimension_mismatches() {
4677 let vectors = vec![1.0f32; 6];
4678 let raw_f16 = vec![0u8; 12];
4679 let raw_u8 = vec![0u8; 6];
4680 let mut scores = vec![0.0f32; 2];
4681
4682 for invalid_query in [vec![1.0, 2.0], vec![1.0, 2.0, 3.0, 4.0]] {
4683 assert!(
4684 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
4685 batch_cosine_scores(&invalid_query, &vectors, 3, &mut scores)
4686 }))
4687 .is_err()
4688 );
4689 assert!(
4690 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
4691 batch_dot_scores_f16(&invalid_query, &raw_f16, 3, &mut scores)
4692 }))
4693 .is_err()
4694 );
4695 assert!(
4696 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
4697 batch_cosine_scores_u8(&invalid_query, &raw_u8, 3, &mut scores)
4698 }))
4699 .is_err()
4700 );
4701 }
4702 }
4703
4704 #[test]
4705 fn vector_simd_boundaries_reject_truncated_storage() {
4706 let query = [1.0f32, 2.0, 3.0];
4707 let mut scores = [0.0f32; 2];
4708
4709 assert!(
4710 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
4711 batch_dot_scores(&query, &[0.0; 5], 3, &mut scores)
4712 }))
4713 .is_err()
4714 );
4715 assert!(
4716 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
4717 batch_cosine_scores_f16(&query, &[0u8; 11], 3, &mut scores)
4718 }))
4719 .is_err()
4720 );
4721 assert!(
4722 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
4723 dot_product_f32(&query, &query, 4)
4724 }))
4725 .is_err()
4726 );
4727 }
4728
4729 #[test]
4730 fn test_unpack_8bit() {
4731 let input: Vec<u8> = (0..128).collect();
4732 let mut output = vec![0u32; 128];
4733 unpack_8bit(&input, &mut output, 128);
4734
4735 for (i, &v) in output.iter().enumerate() {
4736 assert_eq!(v, i as u32);
4737 }
4738 }
4739
4740 #[test]
4741 fn test_unpack_16bit() {
4742 let mut input = vec![0u8; 256];
4743 for i in 0..128 {
4744 let val = (i * 100) as u16;
4745 input[i * 2] = val as u8;
4746 input[i * 2 + 1] = (val >> 8) as u8;
4747 }
4748
4749 let mut output = vec![0u32; 128];
4750 unpack_16bit(&input, &mut output, 128);
4751
4752 for (i, &v) in output.iter().enumerate() {
4753 assert_eq!(v, (i * 100) as u32);
4754 }
4755 }
4756
4757 #[test]
4758 fn test_unpack_32bit() {
4759 let mut input = vec![0u8; 512];
4760 for i in 0..128 {
4761 let val = (i * 1000) as u32;
4762 let bytes = val.to_le_bytes();
4763 input[i * 4..i * 4 + 4].copy_from_slice(&bytes);
4764 }
4765
4766 let mut output = vec![0u32; 128];
4767 unpack_32bit(&input, &mut output, 128);
4768
4769 for (i, &v) in output.iter().enumerate() {
4770 assert_eq!(v, (i * 1000) as u32);
4771 }
4772 }
4773
4774 #[test]
4775 fn test_delta_decode() {
4776 let deltas = vec![4u32, 4, 9, 19];
4780 let mut output = vec![0u32; 5];
4781
4782 delta_decode(&mut output, &deltas, 10, 5);
4783
4784 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4785 }
4786
4787 #[test]
4788 fn test_add_one() {
4789 let mut values = vec![0u32, 1, 2, 3, 4, 5, 6, 7];
4790 add_one(&mut values, 8);
4791
4792 assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
4793 }
4794
4795 #[test]
4796 fn test_bits_needed() {
4797 assert_eq!(bits_needed(0), 0);
4798 assert_eq!(bits_needed(1), 1);
4799 assert_eq!(bits_needed(2), 2);
4800 assert_eq!(bits_needed(3), 2);
4801 assert_eq!(bits_needed(4), 3);
4802 assert_eq!(bits_needed(255), 8);
4803 assert_eq!(bits_needed(256), 9);
4804 assert_eq!(bits_needed(u32::MAX), 32);
4805 }
4806
4807 #[test]
4808 fn test_unpack_8bit_delta_decode() {
4809 let input: Vec<u8> = vec![4, 4, 9, 19];
4813 let mut output = vec![0u32; 5];
4814
4815 unpack_8bit_delta_decode(&input, &mut output, 10, 5);
4816
4817 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4818 }
4819
4820 #[test]
4821 fn test_unpack_16bit_delta_decode() {
4822 let mut input = vec![0u8; 8];
4826 for (i, &delta) in [499u16, 499, 999, 1999].iter().enumerate() {
4827 input[i * 2] = delta as u8;
4828 input[i * 2 + 1] = (delta >> 8) as u8;
4829 }
4830 let mut output = vec![0u32; 5];
4831
4832 unpack_16bit_delta_decode(&input, &mut output, 100, 5);
4833
4834 assert_eq!(output, vec![100, 600, 1100, 2100, 4100]);
4835 }
4836
4837 #[test]
4838 fn test_fused_vs_separate_8bit() {
4839 let input: Vec<u8> = (0..127).collect();
4841 let first_value = 1000u32;
4842 let count = 128;
4843
4844 let mut unpacked = vec![0u32; 128];
4846 unpack_8bit(&input, &mut unpacked, 127);
4847 let mut separate_output = vec![0u32; 128];
4848 delta_decode(&mut separate_output, &unpacked, first_value, count);
4849
4850 let mut fused_output = vec![0u32; 128];
4852 unpack_8bit_delta_decode(&input, &mut fused_output, first_value, count);
4853
4854 assert_eq!(separate_output, fused_output);
4855 }
4856
4857 #[test]
4858 fn test_round_bit_width() {
4859 assert_eq!(round_bit_width(0), 0);
4860 assert_eq!(round_bit_width(1), 8);
4861 assert_eq!(round_bit_width(5), 8);
4862 assert_eq!(round_bit_width(8), 8);
4863 assert_eq!(round_bit_width(9), 16);
4864 assert_eq!(round_bit_width(12), 16);
4865 assert_eq!(round_bit_width(16), 16);
4866 assert_eq!(round_bit_width(17), 32);
4867 assert_eq!(round_bit_width(24), 32);
4868 assert_eq!(round_bit_width(32), 32);
4869 }
4870
4871 #[test]
4872 fn test_rounded_bitwidth_from_exact() {
4873 assert_eq!(RoundedBitWidth::from_exact(0), RoundedBitWidth::Zero);
4874 assert_eq!(RoundedBitWidth::from_exact(1), RoundedBitWidth::Bits8);
4875 assert_eq!(RoundedBitWidth::from_exact(8), RoundedBitWidth::Bits8);
4876 assert_eq!(RoundedBitWidth::from_exact(9), RoundedBitWidth::Bits16);
4877 assert_eq!(RoundedBitWidth::from_exact(16), RoundedBitWidth::Bits16);
4878 assert_eq!(RoundedBitWidth::from_exact(17), RoundedBitWidth::Bits32);
4879 assert_eq!(RoundedBitWidth::from_exact(32), RoundedBitWidth::Bits32);
4880 }
4881
4882 #[test]
4883 fn test_pack_unpack_rounded_8bit() {
4884 let values: Vec<u32> = (0..128).map(|i| i % 256).collect();
4885 let mut packed = vec![0u8; 128];
4886
4887 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits8, &mut packed);
4888 assert_eq!(bytes_written, 128);
4889
4890 let mut unpacked = vec![0u32; 128];
4891 unpack_rounded(&packed, RoundedBitWidth::Bits8, &mut unpacked, 128);
4892
4893 assert_eq!(values, unpacked);
4894 }
4895
4896 #[test]
4897 fn test_pack_unpack_rounded_16bit() {
4898 let values: Vec<u32> = (0..128).map(|i| i * 100).collect();
4899 let mut packed = vec![0u8; 256];
4900
4901 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits16, &mut packed);
4902 assert_eq!(bytes_written, 256);
4903
4904 let mut unpacked = vec![0u32; 128];
4905 unpack_rounded(&packed, RoundedBitWidth::Bits16, &mut unpacked, 128);
4906
4907 assert_eq!(values, unpacked);
4908 }
4909
4910 #[test]
4911 fn test_pack_unpack_rounded_32bit() {
4912 let values: Vec<u32> = (0..128).map(|i| i * 100000).collect();
4913 let mut packed = vec![0u8; 512];
4914
4915 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits32, &mut packed);
4916 assert_eq!(bytes_written, 512);
4917
4918 let mut unpacked = vec![0u32; 128];
4919 unpack_rounded(&packed, RoundedBitWidth::Bits32, &mut unpacked, 128);
4920
4921 assert_eq!(values, unpacked);
4922 }
4923
4924 #[test]
4925 fn test_unpack_rounded_delta_decode() {
4926 let input: Vec<u8> = vec![4, 4, 9, 19];
4931 let mut output = vec![0u32; 5];
4932
4933 unpack_rounded_delta_decode(&input, RoundedBitWidth::Bits8, &mut output, 10, 5);
4934
4935 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4936 }
4937
4938 #[test]
4939 fn test_unpack_rounded_delta_decode_zero() {
4940 let input: Vec<u8> = vec![];
4942 let mut output = vec![0u32; 5];
4943
4944 unpack_rounded_delta_decode(&input, RoundedBitWidth::Zero, &mut output, 100, 5);
4945
4946 assert_eq!(output, vec![100, 101, 102, 103, 104]);
4947 }
4948
4949 #[test]
4954 fn test_dequantize_uint8() {
4955 let input: Vec<u8> = vec![0, 128, 255, 64, 192];
4956 let mut output = vec![0.0f32; 5];
4957 let scale = 0.1;
4958 let min_val = 1.0;
4959
4960 dequantize_uint8(&input, &mut output, scale, min_val, 5);
4961
4962 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); }
4969
4970 #[test]
4971 fn test_dequantize_uint8_large() {
4972 let input: Vec<u8> = (0..128).collect();
4974 let mut output = vec![0.0f32; 128];
4975 let scale = 2.0;
4976 let min_val = -10.0;
4977
4978 dequantize_uint8(&input, &mut output, scale, min_val, 128);
4979
4980 for (i, &out) in output.iter().enumerate().take(128) {
4981 let expected = i as f32 * scale + min_val;
4982 assert!(
4983 (out - expected).abs() < 1e-5,
4984 "Mismatch at {}: expected {}, got {}",
4985 i,
4986 expected,
4987 out
4988 );
4989 }
4990 }
4991
4992 #[test]
4993 fn test_dot_product_f32() {
4994 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
4995 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0];
4996
4997 let result = dot_product_f32(&a, &b, 5);
4998
4999 assert!((result - 70.0).abs() < 1e-5);
5001 }
5002
5003 #[test]
5004 fn test_dot_product_f32_large() {
5005 let a: Vec<f32> = (0..128).map(|i| i as f32).collect();
5007 let b: Vec<f32> = (0..128).map(|i| (i + 1) as f32).collect();
5008
5009 let result = dot_product_f32(&a, &b, 128);
5010
5011 let expected: f32 = (0..128).map(|i| (i as f32) * ((i + 1) as f32)).sum();
5013 assert!(
5014 (result - expected).abs() < 1e-3,
5015 "Expected {}, got {}",
5016 expected,
5017 result
5018 );
5019 }
5020
5021 #[test]
5022 fn test_fused_dot_norm() {
5023 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
5024 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
5025 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
5026
5027 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
5028 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
5029 assert!(
5030 (dot - expected_dot).abs() < 1e-5,
5031 "dot: expected {}, got {}",
5032 expected_dot,
5033 dot
5034 );
5035 assert!(
5036 (norm_b - expected_norm).abs() < 1e-5,
5037 "norm: expected {}, got {}",
5038 expected_norm,
5039 norm_b
5040 );
5041 }
5042
5043 #[test]
5044 fn test_fused_dot_norm_large() {
5045 let a: Vec<f32> = (0..768).map(|i| (i as f32) * 0.01).collect();
5046 let b: Vec<f32> = (0..768).map(|i| (i as f32) * 0.02 + 0.5).collect();
5047 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
5048
5049 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
5050 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
5051 assert!(
5052 (dot - expected_dot).abs() < 1.0,
5053 "dot: expected {}, got {}",
5054 expected_dot,
5055 dot
5056 );
5057 assert!(
5058 (norm_b - expected_norm).abs() < 1.0,
5059 "norm: expected {}, got {}",
5060 expected_norm,
5061 norm_b
5062 );
5063 }
5064
5065 #[test]
5066 fn test_batch_cosine_scores() {
5067 let query = vec![1.0f32, 0.0, 0.0];
5069 let vectors = vec![
5070 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, -1.0, 0.0, 0.0, 0.5, 0.5, 0.0, ];
5075 let mut scores = vec![0f32; 4];
5076 batch_cosine_scores(&query, &vectors, 3, &mut scores);
5077
5078 assert!((scores[0] - 1.0).abs() < 1e-5, "identical: {}", scores[0]);
5079 assert!(scores[1].abs() < 1e-5, "orthogonal: {}", scores[1]);
5080 assert!((scores[2] - (-1.0)).abs() < 1e-5, "opposite: {}", scores[2]);
5081 let expected_45 = 0.5f32 / (0.5f32.powi(2) + 0.5f32.powi(2)).sqrt();
5082 assert!(
5083 (scores[3] - expected_45).abs() < 1e-5,
5084 "45deg: expected {}, got {}",
5085 expected_45,
5086 scores[3]
5087 );
5088 }
5089
5090 #[test]
5091 fn test_batch_cosine_scores_matches_individual() {
5092 let query: Vec<f32> = (0..128).map(|i| (i as f32) * 0.1).collect();
5093 let n = 50;
5094 let dim = 128;
5095 let vectors: Vec<f32> = (0..n * dim).map(|i| ((i * 7 + 3) as f32) * 0.01).collect();
5096
5097 let mut batch_scores = vec![0f32; n];
5098 batch_cosine_scores(&query, &vectors, dim, &mut batch_scores);
5099
5100 for i in 0..n {
5101 let vec_i = &vectors[i * dim..(i + 1) * dim];
5102 let individual = cosine_similarity(&query, vec_i);
5103 assert!(
5104 (batch_scores[i] - individual).abs() < 1e-5,
5105 "vec {}: batch={}, individual={}",
5106 i,
5107 batch_scores[i],
5108 individual
5109 );
5110 }
5111 }
5112
5113 #[test]
5114 fn test_batch_cosine_scores_empty() {
5115 let query = vec![1.0f32, 2.0, 3.0];
5116 let vectors: Vec<f32> = vec![];
5117 let mut scores: Vec<f32> = vec![];
5118 batch_cosine_scores(&query, &vectors, 3, &mut scores);
5119 assert!(scores.is_empty());
5120 }
5121
5122 #[test]
5123 fn test_batch_cosine_scores_zero_query() {
5124 let query = vec![0.0f32, 0.0, 0.0];
5125 let vectors = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
5126 let mut scores = vec![0f32; 2];
5127 batch_cosine_scores(&query, &vectors, 3, &mut scores);
5128 assert_eq!(scores[0], 0.0);
5129 assert_eq!(scores[1], 0.0);
5130 }
5131
5132 #[test]
5137 fn test_f16_roundtrip_normal() {
5138 for &v in &[0.0f32, 1.0, -1.0, 0.5, -0.5, 0.333, 65504.0] {
5139 let h = f32_to_f16(v);
5140 let back = f16_to_f32(h);
5141 let err = (back - v).abs() / v.abs().max(1e-6);
5142 assert!(
5143 err < 0.002,
5144 "f16 roundtrip {v} → {h:#06x} → {back}, rel err {err}"
5145 );
5146 }
5147 }
5148
5149 #[test]
5150 fn test_f16_special() {
5151 assert_eq!(f16_to_f32(f32_to_f16(0.0)), 0.0);
5153 assert_eq!(f32_to_f16(-0.0), 0x8000);
5155 assert!(f16_to_f32(f32_to_f16(f32::INFINITY)).is_infinite());
5157 assert!(f16_to_f32(f32_to_f16(f32::NAN)).is_nan());
5159 }
5160
5161 #[test]
5162 fn test_f16_embedding_range() {
5163 let values: Vec<f32> = (-100..=100).map(|i| i as f32 / 100.0).collect();
5165 for &v in &values {
5166 let back = f16_to_f32(f32_to_f16(v));
5167 assert!((back - v).abs() < 0.001, "f16 error for {v}: got {back}");
5168 }
5169 }
5170
5171 #[test]
5176 fn test_u8_roundtrip() {
5177 assert_eq!(f32_to_u8_saturating(-1.0), 0);
5179 assert_eq!(f32_to_u8_saturating(1.0), 255);
5180 assert_eq!(f32_to_u8_saturating(0.0), 127); assert_eq!(f32_to_u8_saturating(-2.0), 0);
5184 assert_eq!(f32_to_u8_saturating(2.0), 255);
5185 }
5186
5187 #[test]
5188 fn test_u8_dequantize() {
5189 assert!((u8_to_f32(0) - (-1.0)).abs() < 0.01);
5190 assert!((u8_to_f32(255) - 1.0).abs() < 0.01);
5191 assert!((u8_to_f32(127) - 0.0).abs() < 0.01);
5192 }
5193
5194 #[test]
5199 fn test_batch_cosine_scores_f16() {
5200 let query = vec![0.6f32, 0.8, 0.0, 0.0];
5201 let dim = 4;
5202 let vecs_f32 = vec![
5203 0.6f32, 0.8, 0.0, 0.0, 0.0, 0.0, 0.6, 0.8, ];
5206
5207 let mut f16_buf = vec![0u16; 8];
5209 batch_f32_to_f16(&vecs_f32, &mut f16_buf);
5210 let raw: &[u8] =
5211 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
5212
5213 let mut scores = vec![0f32; 2];
5214 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
5215
5216 assert!(
5217 (scores[0] - 1.0).abs() < 0.01,
5218 "identical vectors: {}",
5219 scores[0]
5220 );
5221 assert!(scores[1].abs() < 0.01, "orthogonal vectors: {}", scores[1]);
5222 }
5223
5224 #[test]
5225 fn test_batch_cosine_scores_u8() {
5226 let query = vec![0.6f32, 0.8, 0.0, 0.0];
5227 let dim = 4;
5228 let vecs_f32 = vec![
5229 0.6f32, 0.8, 0.0, 0.0, -0.6, -0.8, 0.0, 0.0, ];
5232
5233 let mut u8_buf = vec![0u8; 8];
5235 batch_f32_to_u8(&vecs_f32, &mut u8_buf);
5236
5237 let mut scores = vec![0f32; 2];
5238 batch_cosine_scores_u8(&query, &u8_buf, dim, &mut scores);
5239
5240 assert!(scores[0] > 0.95, "similar vectors: {}", scores[0]);
5241 assert!(scores[1] < -0.95, "opposite vectors: {}", scores[1]);
5242 }
5243
5244 #[test]
5245 fn test_batch_cosine_scores_f16_large_dim() {
5246 let dim = 768;
5248 let query: Vec<f32> = (0..dim).map(|i| (i as f32 / dim as f32) - 0.5).collect();
5249 let vec2: Vec<f32> = query.iter().map(|x| x * 0.9 + 0.01).collect();
5250
5251 let mut all_vecs = query.clone();
5252 all_vecs.extend_from_slice(&vec2);
5253
5254 let mut f16_buf = vec![0u16; all_vecs.len()];
5255 batch_f32_to_f16(&all_vecs, &mut f16_buf);
5256 let raw: &[u8] =
5257 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
5258
5259 let mut scores = vec![0f32; 2];
5260 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
5261
5262 assert!((scores[0] - 1.0).abs() < 0.01, "self-sim: {}", scores[0]);
5264 assert!(scores[1] > 0.99, "scaled-sim: {}", scores[1]);
5266 }
5267
5268 #[test]
5273 fn test_hamming_distance_identical() {
5274 let a = vec![0xAA; 64];
5275 assert_eq!(hamming_distance(&a, &a), 0);
5276 }
5277
5278 #[test]
5279 fn test_hamming_distance_opposite() {
5280 let a = vec![0xFF; 32];
5281 let b = vec![0x00; 32];
5282 assert_eq!(hamming_distance(&a, &b), 256);
5283 }
5284
5285 #[test]
5286 fn test_hamming_distance_known() {
5287 let a = vec![0xAA];
5289 let b = vec![0x55];
5290 assert_eq!(hamming_distance(&a, &b), 8);
5291
5292 let a = vec![0xFF, 0x00];
5294 let b = vec![0x00, 0x00];
5295 assert_eq!(hamming_distance(&a, &b), 8);
5296 }
5297
5298 #[test]
5299 fn test_hamming_distance_single_bit() {
5300 let a = vec![0x00; 16];
5301 let mut b = vec![0x00; 16];
5302 b[7] = 0x01; assert_eq!(hamming_distance(&a, &b), 1);
5304 }
5305
5306 #[test]
5307 fn test_hamming_distance_empty() {
5308 let a: Vec<u8> = vec![];
5309 assert_eq!(hamming_distance(&a, &a), 0);
5310 }
5311
5312 #[test]
5313 fn test_hamming_distance_remainder_path() {
5314 let a = vec![0xFF; 17];
5316 let b = vec![0x00; 17];
5317 assert_eq!(hamming_distance(&a, &b), 136); let a = vec![0xFF; 33];
5321 let b = vec![0x00; 33];
5322 assert_eq!(hamming_distance(&a, &b), 264); }
5324
5325 #[test]
5326 fn test_hamming_distance_large() {
5327 let a = vec![0xFF; 4096];
5329 let b = vec![0x00; 4096];
5330 assert_eq!(hamming_distance(&a, &b), 32768);
5331 }
5332
5333 #[test]
5334 fn test_hamming_distance_scalar_matches() {
5335 for size in [1, 7, 8, 15, 16, 31, 32, 63, 64, 100, 128, 255, 256] {
5337 let a: Vec<u8> = (0..size).map(|i| (i * 37 + 13) as u8).collect();
5338 let b: Vec<u8> = (0..size).map(|i| (i * 53 + 7) as u8).collect();
5339 let expected = hamming_distance_scalar(&a, &b);
5340 let got = hamming_distance(&a, &b);
5341 assert_eq!(got, expected, "mismatch at size {size}");
5342 }
5343 }
5344
5345 #[test]
5350 fn test_batch_hamming_scores_identical() {
5351 let query = vec![0xAA; 16];
5352 let db = vec![0xAA; 16]; let mut scores = vec![0f32; 1];
5354 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
5355 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
5356 }
5357
5358 #[test]
5359 fn test_batch_hamming_scores_opposite() {
5360 let query = vec![0xFF; 16];
5361 let db = vec![0x00; 16];
5362 let mut scores = vec![0f32; 1];
5363 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
5364 assert!((scores[0] - 0.0).abs() < 1e-6, "opposite: {}", scores[0]);
5365 }
5366
5367 #[test]
5368 fn test_batch_hamming_scores_multiple() {
5369 let byte_len = 8;
5370 let dim_bits = 64;
5371 let query = vec![0xFF; byte_len];
5372 let mut db = Vec::new();
5373 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];
5378 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut scores);
5379
5380 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
5381 assert!((scores[1] - 0.0).abs() < 1e-6, "opposite: {}", scores[1]);
5382 assert!((scores[2] - 0.5).abs() < 1e-6, "half: {}", scores[2]);
5383 }
5384
5385 #[test]
5386 fn test_batch_hamming_scores_empty() {
5387 let query = vec![0xFF; 8];
5388 let db: Vec<u8> = vec![];
5389 let mut scores: Vec<f32> = vec![];
5390 batch_hamming_scores(&query, &db, 8, 64, &mut scores);
5391 assert!(scores.is_empty());
5392 }
5393
5394 #[test]
5395 fn test_batch_hamming_scores_zero_byte_len() {
5396 let query: Vec<u8> = vec![];
5397 let db: Vec<u8> = vec![];
5398 let mut scores = vec![0f32; 1];
5399 batch_hamming_scores(&query, &db, 0, 0, &mut scores);
5400 assert_eq!(scores[0], 0.0);
5402 }
5403
5404 fn hamming_matrix(rows: usize, byte_len: usize) -> (Vec<u8>, Vec<u8>) {
5409 let query: Vec<u8> = (0..byte_len).map(|i| (i * 31 + 5) as u8).collect();
5410 let db: Vec<u8> = (0..rows * byte_len)
5411 .map(|i| (i * 97 + i / byte_len * 11 + 3) as u8)
5412 .collect();
5413 (query, db)
5414 }
5415
5416 #[test]
5419 fn batched_hamming_distances_match_scalar_for_every_row_count() {
5420 let kernels = [HammingKernel::resolve(), HammingKernel::Scalar];
5421 for byte_len in [1, 7, 8, 15, 16, 31, 32, 33, 63, 64, 65, 128, 320] {
5424 for rows in [1, 2, 3, 4, 5, 7, 8, 9, 64, 70] {
5425 let (query, db) = hamming_matrix(rows, byte_len);
5426 let mut got = vec![0u32; rows];
5427 for kernel in kernels {
5428 kernel.distances(&query, &db, byte_len, &mut got);
5429 for (row, &distance) in got.iter().enumerate() {
5430 let expected = hamming_distance_scalar(
5431 &query,
5432 &db[row * byte_len..(row + 1) * byte_len],
5433 );
5434 assert_eq!(
5435 distance, expected,
5436 "{kernel:?}: row {row} of {rows} at byte_len {byte_len}"
5437 );
5438 }
5439 }
5440 }
5441 }
5442 }
5443
5444 #[test]
5445 fn gathered_hamming_distances_follow_row_ids() {
5446 let kernel = HammingKernel::resolve();
5447 let byte_len = 320;
5448 let rows = 37;
5449 let (query, db) = hamming_matrix(rows, byte_len);
5450 let ids: Vec<u32> = [36, 0, 17, 17, 5, 31, 2, 9, 9, 36, 1].into_iter().collect();
5453 let mut got = vec![0u32; ids.len()];
5454 kernel.gather_distances(&query, &db, byte_len, &ids, &mut got);
5455 for (slot, &id) in ids.iter().enumerate() {
5456 let start = id as usize * byte_len;
5457 let expected = hamming_distance_scalar(&query, &db[start..start + byte_len]);
5458 assert_eq!(got[slot], expected, "slot {slot} for row {id}");
5459 }
5460 }
5461
5462 #[test]
5463 fn resolved_kernel_matches_scalar_pairwise() {
5464 let kernel = HammingKernel::resolve();
5465 for byte_len in [1, 8, 32, 64, 65, 320, 4096] {
5466 let (query, db) = hamming_matrix(1, byte_len);
5467 assert_eq!(
5468 kernel.distance(&query, &db),
5469 hamming_distance_scalar(&query, &db),
5470 "byte_len {byte_len}"
5471 );
5472 }
5473 }
5474
5475 #[test]
5478 fn hamming_kernel_routes_sub_vector_codes_to_the_scalar_loop() {
5479 let kernel = HammingKernel::resolve();
5480 let width = kernel.vector_bytes();
5481 assert_eq!(HammingKernel::Scalar.for_byte_len(1), HammingKernel::Scalar);
5482 assert_eq!(kernel.for_byte_len(width - 1), HammingKernel::Scalar);
5483 assert_eq!(kernel.for_byte_len(width), kernel);
5484 assert_eq!(kernel.for_byte_len(width * 5 + 3), kernel);
5485 for byte_len in [width - 1, width, width + 1] {
5487 let (query, db) = hamming_matrix(9, byte_len);
5488 let mut got = vec![0u32; 9];
5489 kernel.distances(&query, &db, byte_len, &mut got);
5490 for (row, &distance) in got.iter().enumerate() {
5491 assert_eq!(
5492 distance,
5493 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]),
5494 "row {row} at byte_len {byte_len}"
5495 );
5496 }
5497 }
5498 }
5499
5500 #[test]
5501 fn scores_from_hamming_matches_batch_scores_across_blocks() {
5502 let kernel = HammingKernel::resolve();
5503 let byte_len = 320;
5504 let dim_bits = byte_len * 8;
5505 let rows = HAMMING_DISTANCE_BLOCK * 2 + 3;
5507 let (query, db) = hamming_matrix(rows, byte_len);
5508 let mut expected = vec![0f32; rows];
5509 for (row, score) in expected.iter_mut().enumerate() {
5510 let distance =
5511 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
5512 *score = 1.0 - distance as f32 / dim_bits as f32;
5513 }
5514 let mut got = vec![0f32; rows];
5515 scores_from_hamming(kernel, &query, &db, byte_len, dim_bits, &mut got);
5516 for (row, (&got, &want)) in got.iter().zip(expected.iter()).enumerate() {
5517 assert!((got - want).abs() < 1e-6, "row {row}: {got} vs {want}");
5518 }
5519 let mut public = vec![0f32; rows];
5520 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut public);
5521 assert_eq!(got, public);
5522 }
5523}
5524
5525#[inline]
5533pub(crate) fn intersect_posting_blocks(
5534 left: &[u32],
5535 a: &mut usize,
5536 right: &[u32],
5537 b: &mut usize,
5538 pairs: &mut [(u8, u8)],
5539) -> usize {
5540 assert!(left.len() <= 128 && right.len() <= 128);
5541 assert!(*a <= left.len() && *b <= right.len());
5542 let (mut left_pos, mut right_pos) = (*a, *b);
5543 let mut count = 0;
5544 while left_pos < left.len() && right_pos < right.len() && count < pairs.len() {
5545 let doc = left[left_pos];
5546 if let Some(group) = right.get(right_pos..right_pos + 8) {
5547 let group: &[u32; 8] = group.try_into().unwrap();
5548 if group[7] < doc {
5549 right_pos += 8;
5550 continue;
5551 }
5552 if doc < group[0] {
5553 left_pos += find_first_ge_u32(&left[left_pos..], group[0]);
5554 continue;
5555 }
5556 if let Some(lane) = equal_lane_8(group, doc) {
5557 pairs[count] = (left_pos as u8, (right_pos + lane) as u8);
5558 count += 1;
5559 left_pos += 1;
5560 if count == pairs.len() {
5561 right_pos += lane + 1;
5562 break;
5563 }
5564 } else {
5565 left_pos += 1;
5566 }
5567 } else {
5568 match doc.cmp(&right[right_pos]) {
5569 std::cmp::Ordering::Less => left_pos += 1,
5570 std::cmp::Ordering::Greater => right_pos += 1,
5571 std::cmp::Ordering::Equal => {
5572 pairs[count] = (left_pos as u8, right_pos as u8);
5573 count += 1;
5574 left_pos += 1;
5575 right_pos += 1;
5576 }
5577 }
5578 }
5579 }
5580 *a = left_pos;
5581 *b = right_pos;
5582 count
5583}
5584
5585#[inline]
5586fn equal_lane_8(values: &[u32; 8], target: u32) -> Option<usize> {
5587 #[cfg(target_arch = "x86_64")]
5588 if avx2::is_available() {
5589 let mask = unsafe { equal_mask_8_avx2(values, target) };
5591 return (mask != 0).then(|| mask.trailing_zeros() as usize);
5592 }
5593 #[cfg(target_arch = "aarch64")]
5594 if neon::is_available() {
5595 let mask = unsafe { equal_mask_8_neon(values, target) };
5597 return (mask != 0).then(|| mask.trailing_zeros() as usize);
5598 }
5599 values.iter().position(|&value| value == target)
5600}
5601
5602#[cfg(target_arch = "x86_64")]
5603#[target_feature(enable = "avx2")]
5604unsafe fn equal_mask_8_avx2(values: &[u32; 8], target: u32) -> u32 {
5605 use std::arch::x86_64::*;
5606 let values = unsafe { _mm256_loadu_si256(values.as_ptr().cast()) };
5608 _mm256_movemask_ps(_mm256_castsi256_ps(_mm256_cmpeq_epi32(
5609 values,
5610 _mm256_set1_epi32(target as i32),
5611 ))) as u32
5612}
5613
5614#[cfg(target_arch = "aarch64")]
5615#[target_feature(enable = "neon")]
5616unsafe fn equal_mask_8_neon(values: &[u32; 8], target: u32) -> u32 {
5617 use std::arch::aarch64::*;
5618 unsafe {
5620 let target = vdupq_n_u32(target);
5621 let weights = [1u32, 2, 4, 8];
5622 let weights = vld1q_u32(weights.as_ptr());
5623 let lo = vceqq_u32(vld1q_u32(values.as_ptr()), target);
5624 let hi = vceqq_u32(vld1q_u32(values.as_ptr().add(4)), target);
5625 vaddvq_u32(vandq_u32(lo, weights)) | (vaddvq_u32(vandq_u32(hi, weights)) << 4)
5626 }
5627}
5628
5629#[inline]
5632pub(crate) fn find_first_ge_block_from(docs: &[u32], from: usize, target: u32) -> usize {
5633 debug_assert!(from <= docs.len());
5634 if let Ok(block) = <&[u32; 128]>::try_from(docs) {
5635 if block[127] < target {
5636 return 128;
5637 }
5638 let mut base = 0;
5639 let mut step = 64;
5640 while step != 0 {
5641 base += usize::from(block[base + step - 1] < target) * step;
5642 step >>= 1;
5643 }
5644 base.max(from)
5645 } else {
5646 from + find_first_ge_u32(&docs[from..], target)
5647 }
5648}
5649
5650#[inline]
5659pub fn find_first_ge_u32(slice: &[u32], target: u32) -> usize {
5660 #[cfg(target_arch = "aarch64")]
5661 {
5662 if neon::is_available() {
5663 return unsafe { find_first_ge_u32_neon(slice, target) };
5665 }
5666 slice.partition_point(|&d| d < target)
5667 }
5668
5669 #[cfg(target_arch = "x86_64")]
5670 {
5671 if avx2::is_available() {
5672 return unsafe { find_first_ge_u32_avx2(slice, target) };
5674 }
5675 unsafe { find_first_ge_u32_sse(slice, target) }
5679 }
5680
5681 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5683 {
5684 slice.partition_point(|&d| d < target)
5685 }
5686}
5687
5688#[cfg(target_arch = "aarch64")]
5689#[target_feature(enable = "neon")]
5690#[allow(unsafe_op_in_unsafe_fn)]
5691unsafe fn find_first_ge_u32_neon(slice: &[u32], target: u32) -> usize {
5692 use std::arch::aarch64::*;
5693
5694 let n = slice.len();
5695 let ptr = slice.as_ptr();
5696 let target_vec = vdupq_n_u32(target);
5697 let bit_mask: uint32x4_t = core::mem::transmute([1u32, 2u32, 4u32, 8u32]);
5699
5700 let chunks = n / 16;
5701 let mut base = 0usize;
5702
5703 for _ in 0..chunks {
5705 let v0 = vld1q_u32(ptr.add(base));
5706 let v1 = vld1q_u32(ptr.add(base + 4));
5707 let v2 = vld1q_u32(ptr.add(base + 8));
5708 let v3 = vld1q_u32(ptr.add(base + 12));
5709
5710 let c0 = vcgeq_u32(v0, target_vec);
5711 let c1 = vcgeq_u32(v1, target_vec);
5712 let c2 = vcgeq_u32(v2, target_vec);
5713 let c3 = vcgeq_u32(v3, target_vec);
5714
5715 let m0 = vaddvq_u32(vandq_u32(c0, bit_mask));
5716 if m0 != 0 {
5717 return base + m0.trailing_zeros() as usize;
5718 }
5719 let m1 = vaddvq_u32(vandq_u32(c1, bit_mask));
5720 if m1 != 0 {
5721 return base + 4 + m1.trailing_zeros() as usize;
5722 }
5723 let m2 = vaddvq_u32(vandq_u32(c2, bit_mask));
5724 if m2 != 0 {
5725 return base + 8 + m2.trailing_zeros() as usize;
5726 }
5727 let m3 = vaddvq_u32(vandq_u32(c3, bit_mask));
5728 if m3 != 0 {
5729 return base + 12 + m3.trailing_zeros() as usize;
5730 }
5731 base += 16;
5732 }
5733
5734 while base + 4 <= n {
5736 let vals = vld1q_u32(ptr.add(base));
5737 let cmp = vcgeq_u32(vals, target_vec);
5738 let mask = vaddvq_u32(vandq_u32(cmp, bit_mask));
5739 if mask != 0 {
5740 return base + mask.trailing_zeros() as usize;
5741 }
5742 base += 4;
5743 }
5744
5745 while base < n {
5747 if *slice.get_unchecked(base) >= target {
5748 return base;
5749 }
5750 base += 1;
5751 }
5752 n
5753}
5754
5755#[cfg(target_arch = "x86_64")]
5756#[target_feature(enable = "sse2")]
5757#[allow(unsafe_op_in_unsafe_fn)]
5758unsafe fn find_first_ge_u32_sse(slice: &[u32], target: u32) -> usize {
5759 use std::arch::x86_64::*;
5760
5761 let n = slice.len();
5762 let ptr = slice.as_ptr();
5763
5764 let sign_flip = _mm_set1_epi32(i32::MIN);
5766 let target_xor = _mm_xor_si128(_mm_set1_epi32(target as i32), sign_flip);
5767
5768 let chunks = n / 16;
5769 let mut base = 0usize;
5770
5771 for _ in 0..chunks {
5773 let v0 = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
5774 let v1 = _mm_xor_si128(
5775 _mm_loadu_si128(ptr.add(base + 4) as *const __m128i),
5776 sign_flip,
5777 );
5778 let v2 = _mm_xor_si128(
5779 _mm_loadu_si128(ptr.add(base + 8) as *const __m128i),
5780 sign_flip,
5781 );
5782 let v3 = _mm_xor_si128(
5783 _mm_loadu_si128(ptr.add(base + 12) as *const __m128i),
5784 sign_flip,
5785 );
5786
5787 let ge0 = _mm_or_si128(
5789 _mm_cmpeq_epi32(v0, target_xor),
5790 _mm_cmpgt_epi32(v0, target_xor),
5791 );
5792 let m0 = _mm_movemask_ps(_mm_castsi128_ps(ge0)) as u32;
5793 if m0 != 0 {
5794 return base + m0.trailing_zeros() as usize;
5795 }
5796
5797 let ge1 = _mm_or_si128(
5798 _mm_cmpeq_epi32(v1, target_xor),
5799 _mm_cmpgt_epi32(v1, target_xor),
5800 );
5801 let m1 = _mm_movemask_ps(_mm_castsi128_ps(ge1)) as u32;
5802 if m1 != 0 {
5803 return base + 4 + m1.trailing_zeros() as usize;
5804 }
5805
5806 let ge2 = _mm_or_si128(
5807 _mm_cmpeq_epi32(v2, target_xor),
5808 _mm_cmpgt_epi32(v2, target_xor),
5809 );
5810 let m2 = _mm_movemask_ps(_mm_castsi128_ps(ge2)) as u32;
5811 if m2 != 0 {
5812 return base + 8 + m2.trailing_zeros() as usize;
5813 }
5814
5815 let ge3 = _mm_or_si128(
5816 _mm_cmpeq_epi32(v3, target_xor),
5817 _mm_cmpgt_epi32(v3, target_xor),
5818 );
5819 let m3 = _mm_movemask_ps(_mm_castsi128_ps(ge3)) as u32;
5820 if m3 != 0 {
5821 return base + 12 + m3.trailing_zeros() as usize;
5822 }
5823 base += 16;
5824 }
5825
5826 while base + 4 <= n {
5828 let vals = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
5829 let ge = _mm_or_si128(
5830 _mm_cmpeq_epi32(vals, target_xor),
5831 _mm_cmpgt_epi32(vals, target_xor),
5832 );
5833 let mask = _mm_movemask_ps(_mm_castsi128_ps(ge)) as u32;
5834 if mask != 0 {
5835 return base + mask.trailing_zeros() as usize;
5836 }
5837 base += 4;
5838 }
5839
5840 while base < n {
5842 if *slice.get_unchecked(base) >= target {
5843 return base;
5844 }
5845 base += 1;
5846 }
5847 n
5848}
5849
5850#[cfg(target_arch = "x86_64")]
5851#[target_feature(enable = "avx2")]
5852#[allow(unsafe_op_in_unsafe_fn)]
5853unsafe fn find_first_ge_u32_avx2(slice: &[u32], target: u32) -> usize {
5854 use std::arch::x86_64::*;
5855 let target_vec = _mm256_set1_epi32(target as i32);
5856 let mut base = 0;
5857 while slice.len() - base >= 8 {
5858 let values = _mm256_loadu_si256(slice.as_ptr().add(base).cast());
5859 let ge = _mm256_cmpeq_epi32(_mm256_min_epu32(values, target_vec), target_vec);
5862 let mask = _mm256_movemask_ps(_mm256_castsi256_ps(ge)) as u32;
5863 if mask != 0 {
5864 return base + mask.trailing_zeros() as usize;
5865 }
5866 base += 8;
5867 }
5868 base + find_first_ge_u32_sse(&slice[base..], target)
5869}
5870
5871#[cfg(test)]
5872mod find_first_ge_tests {
5873 use super::find_first_ge_u32;
5874
5875 #[test]
5876 fn test_find_first_ge_basic() {
5877 let data: Vec<u32> = (0..128).map(|i| i * 3).collect(); assert_eq!(find_first_ge_u32(&data, 0), 0);
5879 assert_eq!(find_first_ge_u32(&data, 1), 1); assert_eq!(find_first_ge_u32(&data, 3), 1);
5881 assert_eq!(find_first_ge_u32(&data, 4), 2); assert_eq!(find_first_ge_u32(&data, 381), 127);
5883 assert_eq!(find_first_ge_u32(&data, 382), 128); }
5885
5886 #[test]
5887 fn test_find_first_ge_matches_partition_point() {
5888 let data: Vec<u32> = vec![1, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55, 60, 65, 70, 75];
5889 for target in 0..80 {
5890 let expected = data.partition_point(|&d| d < target);
5891 let actual = find_first_ge_u32(&data, target);
5892 assert_eq!(actual, expected, "target={}", target);
5893 }
5894 }
5895
5896 #[test]
5897 fn test_find_first_ge_small_slices() {
5898 assert_eq!(find_first_ge_u32(&[], 5), 0);
5900 assert_eq!(find_first_ge_u32(&[10], 5), 0);
5902 assert_eq!(find_first_ge_u32(&[10], 10), 0);
5903 assert_eq!(find_first_ge_u32(&[10], 11), 1);
5904 assert_eq!(find_first_ge_u32(&[2, 4, 6], 5), 2);
5906 }
5907
5908 #[test]
5909 fn test_find_first_ge_full_block() {
5910 let data: Vec<u32> = (100..228).collect();
5912 assert_eq!(find_first_ge_u32(&data, 100), 0);
5913 assert_eq!(find_first_ge_u32(&data, 150), 50);
5914 assert_eq!(find_first_ge_u32(&data, 227), 127);
5915 assert_eq!(find_first_ge_u32(&data, 228), 128);
5916 assert_eq!(find_first_ge_u32(&data, 99), 0);
5917 }
5918
5919 #[test]
5920 fn test_find_first_ge_u32_max() {
5921 let data = vec![u32::MAX - 10, u32::MAX - 5, u32::MAX - 1, u32::MAX];
5923 assert_eq!(find_first_ge_u32(&data, u32::MAX - 10), 0);
5924 assert_eq!(find_first_ge_u32(&data, u32::MAX - 7), 1);
5925 assert_eq!(find_first_ge_u32(&data, u32::MAX), 3);
5926 }
5927
5928 #[test]
5932 fn find_first_ge_matches_partition_point_for_every_length_and_target() {
5933 let mut state = 0x9E37_79B9u32;
5934 let mut next = move || {
5935 state ^= state << 13;
5936 state ^= state >> 17;
5937 state ^= state << 5;
5938 state
5939 };
5940 for n in 0..=140usize {
5941 let mut data: Vec<u32> = Vec::with_capacity(n);
5942 let mut value = next() % 8;
5943 for i in 0..n {
5944 let step = match next() % 5 {
5947 0 | 1 => 0,
5948 2 => 1,
5949 3 => next() % 1000,
5950 _ => 0x4000_0000 + next() % 0x1000_0000,
5951 };
5952 value = value.saturating_add(step);
5953 if i + 3 >= n && n > 8 {
5954 value = u32::MAX;
5955 }
5956 data.push(value);
5957 }
5958 assert!(data.windows(2).all(|w| w[0] <= w[1]));
5959 let mut targets: Vec<u32> =
5960 vec![0, 1, u32::MAX - 1, u32::MAX, i32::MAX as u32, 1 << 31];
5961 for &d in &data {
5962 targets.extend([d.saturating_sub(1), d, d.saturating_add(1)]);
5963 }
5964 for target in targets {
5965 let expected = data.partition_point(|&d| d < target);
5966 assert_eq!(
5967 find_first_ge_u32(&data, target),
5968 expected,
5969 "n={n} target={target} data={data:?}"
5970 );
5971 }
5972 }
5973 }
5974}
5975
5976#[cfg(test)]
5980mod algebraic_reduction_tests {
5981 use super::*;
5982
5983 fn algebraic_test_vector(dim: usize, seed: u64) -> Vec<f32> {
5984 let mut state = seed | 1;
5985 (0..dim)
5986 .map(|_| {
5987 state = state
5988 .wrapping_mul(6364136223846793005)
5989 .wrapping_add(1442695040888963407);
5990 ((state >> 33) as f32 / (1u64 << 30) as f32) - 1.0
5991 })
5992 .collect()
5993 }
5994
5995 const ALGEBRAIC_TEST_DIMS: [usize; 18] = [
5998 0, 1, 3, 4, 7, 8, 15, 16, 17, 64, 100, 127, 200, 300, 384, 768, 1000, 1536,
5999 ];
6000
6001 #[test]
6002 fn test_algebraic_squared_l2_matches_f64_reference() {
6003 for dim in ALGEBRAIC_TEST_DIMS {
6004 let a = algebraic_test_vector(dim, 0x51ed_0001);
6005 let b = algebraic_test_vector(dim, 0x51ed_0002);
6006 let reference: f64 = a
6007 .iter()
6008 .zip(&b)
6009 .map(|(&x, &y)| {
6010 let delta = f64::from(x) - f64::from(y);
6011 delta * delta
6012 })
6013 .sum();
6014 let actual = squared_l2_f32(&a, &b);
6015 let tolerance = (reference * 1e-5).max(1e-6);
6016 assert!(
6017 (f64::from(actual) - reference).abs() <= tolerance,
6018 "dim {dim}: squared_l2_f32 {actual} drifted from f64 reference {reference}"
6019 );
6020 }
6021 }
6022
6023 #[test]
6024 fn test_algebraic_squared_l2_uses_shorter_length() {
6025 let a = [1.0f32, 2.0, 3.0, 4.0];
6026 let b = [1.0f32, 4.0];
6027 assert_eq!(squared_l2_f32(&a, &b), 4.0);
6028 assert_eq!(squared_l2_f32(&b, &a), 4.0);
6029 }
6030
6031 #[test]
6032 fn test_algebraic_norm_squared_matches_f64_reference() {
6033 for dim in ALGEBRAIC_TEST_DIMS {
6034 let v = algebraic_test_vector(dim, 0x51ed_0003);
6035 let reference: f64 = v.iter().map(|&x| f64::from(x) * f64::from(x)).sum();
6036 let actual = norm_squared_f32(&v);
6037 let tolerance = (reference * 1e-5).max(1e-6);
6038 assert!(
6039 (f64::from(actual) - reference).abs() <= tolerance,
6040 "dim {dim}: norm_squared_f32 {actual} drifted from f64 reference {reference}"
6041 );
6042 assert!(
6043 (f64::from(norm_f32(&v)) - reference.sqrt()).abs() <= tolerance.sqrt().max(1e-5)
6044 );
6045 }
6046 }
6047
6048 #[test]
6049 fn test_algebraic_dot_scalar_matches_simd_dispatch() {
6050 for dim in ALGEBRAIC_TEST_DIMS {
6051 let a = algebraic_test_vector(dim, 0x51ed_0004);
6052 let b = algebraic_test_vector(dim, 0x51ed_0005);
6053 let dispatched = dot_product_f32(&a, &b, dim);
6054 let scalar = dot_product_f32_scalar(&a, &b);
6055 let tolerance = (dispatched.abs() * 1e-5).max(1e-5);
6056 assert!(
6057 (dispatched - scalar).abs() <= tolerance,
6058 "dim {dim}: scalar dot {scalar} disagrees with dispatched {dispatched}"
6059 );
6060
6061 let (fused_dot, fused_norm) = fused_dot_norm(&a, &b, dim);
6062 let (scalar_dot, scalar_norm) = fused_dot_norm_scalar(&a, &b);
6063 assert!((fused_dot - scalar_dot).abs() <= tolerance);
6064 assert!((fused_norm - scalar_norm).abs() <= (fused_norm.abs() * 1e-5).max(1e-5));
6065 }
6066 }
6067
6068 #[test]
6073 fn test_simd_tail_dims_match_f64_reference_and_resolved_kernels() {
6074 for dim in ALGEBRAIC_TEST_DIMS {
6075 let a = algebraic_test_vector(dim, 0x51ed_0006);
6076 let b = algebraic_test_vector(dim, 0x51ed_0007);
6077 let reference: f64 = a
6078 .iter()
6079 .zip(&b)
6080 .map(|(&x, &y)| f64::from(x) * f64::from(y))
6081 .sum();
6082 let dispatched = dot_product_f32(&a, &b, dim);
6083 let tolerance = (reference.abs() * 1e-5).max(1e-5);
6084 assert!(
6085 (f64::from(dispatched) - reference).abs() <= tolerance,
6086 "dim {dim}: dot {dispatched} drifted from f64 reference {reference}"
6087 );
6088 let kernel = DenseF32Kernel::resolve();
6089 assert_eq!(kernel.dot(&a, &b, dim).to_bits(), dispatched.to_bits());
6090 let (fused_dot, fused_norm) = fused_dot_norm(&a, &b, dim);
6091 let (kernel_dot, kernel_norm) = kernel.fused_dot_norm(&a, &b, dim);
6092 assert_eq!(kernel_dot.to_bits(), fused_dot.to_bits());
6093 assert_eq!(kernel_norm.to_bits(), fused_norm.to_bits());
6094 let norm_reference: f64 = b.iter().map(|&y| f64::from(y) * f64::from(y)).sum();
6095 assert!(
6096 (f64::from(fused_norm) - norm_reference).abs() <= (norm_reference * 1e-5).max(1e-5),
6097 "dim {dim}: fused norm {fused_norm} drifted from {norm_reference}"
6098 );
6099
6100 let query_f16: Vec<u16> = a.iter().map(|&v| f32_to_f16(v)).collect();
6101 let vec_f16: Vec<u16> = b.iter().map(|&v| f32_to_f16(v)).collect();
6102 let f16_kernel = QuantF16Kernel::resolve();
6103 let (d, n) = fused_dot_norm_f16(&query_f16, &vec_f16, dim);
6104 let (kd, kn) = f16_kernel.fused_dot_norm(&query_f16, &vec_f16, dim);
6105 assert_eq!((kd.to_bits(), kn.to_bits()), (d.to_bits(), n.to_bits()));
6106 assert_eq!(
6107 f16_kernel.dot(&query_f16, &vec_f16, dim).to_bits(),
6108 dot_product_f16_quant(&query_f16, &vec_f16, dim).to_bits()
6109 );
6110
6111 let vec_u8: Vec<u8> = b.iter().map(|&v| f32_to_u8_saturating(v)).collect();
6112 let u8_kernel = QuantU8Kernel::resolve();
6113 let (d, n) = fused_dot_norm_u8(&a, &vec_u8, dim);
6114 let (kd, kn) = u8_kernel.fused_dot_norm(&a, &vec_u8, dim);
6115 assert_eq!((kd.to_bits(), kn.to_bits()), (d.to_bits(), n.to_bits()));
6116 assert_eq!(
6117 u8_kernel.dot(&a, &vec_u8, dim).to_bits(),
6118 dot_product_u8_quant(&a, &vec_u8, dim).to_bits()
6119 );
6120 }
6121 }
6122
6123 #[test]
6128 fn test_algebraic_reductions_propagate_non_finite() {
6129 let finite = vec![1.0f32; 8];
6130
6131 let mut with_nan = finite.clone();
6132 with_nan[5] = f32::NAN;
6133 assert!(norm_squared_f32(&with_nan).is_nan());
6134 assert!(squared_l2_f32(&with_nan, &finite).is_nan());
6135 assert!(dot_product_f32_scalar(&with_nan, &finite).is_nan());
6136
6137 let mut with_inf = finite.clone();
6138 with_inf[2] = f32::INFINITY;
6139 assert!(norm_squared_f32(&with_inf).is_infinite());
6140 assert!(squared_l2_f32(&with_inf, &finite).is_infinite());
6141 assert!(dot_product_f32_scalar(&with_inf, &finite).is_infinite());
6142 }
6143}