Skip to main content

hermes_core/structures/
simd.rs

1//! Shared SIMD-accelerated functions for posting list compression
2//!
3//! This module provides platform-optimized implementations for common operations:
4//! - **Unpacking**: Convert packed 8/16/32-bit values to u32 arrays
5//! - **Delta decoding**: Prefix sum for converting deltas to absolute values
6//! - **Add one**: Increment all values in an array (for TF decoding)
7//! - **Float reductions**: dot product, squared L2, squared norm
8//!
9//! Supports:
10//! - **NEON** on aarch64 (Apple Silicon, ARM servers)
11//! - **SSE/SSE4.1** on x86_64 (Intel/AMD)
12//! - **Scalar fallback** for other architectures (algebraic float ops, see below)
13
14// ============================================================================
15// NEON intrinsics for aarch64 (Apple Silicon, ARM servers)
16// ============================================================================
17
18#[cfg(target_arch = "aarch64")]
19#[allow(unsafe_op_in_unsafe_fn)]
20mod neon {
21    use std::arch::aarch64::*;
22
23    /// SIMD unpack for 8-bit values using NEON
24    #[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            // Load 16 bytes
34            let bytes = vld1q_u8(in_ptr);
35
36            // Widen u8 -> u16 -> u32
37            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        // Handle remainder
56        let base = chunks * 16;
57        for i in 0..remainder {
58            output[base + i] = input[base + i] as u32;
59        }
60    }
61
62    /// SIMD unpack for 16-bit values using NEON
63    #[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        // Handle remainder
82        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    /// SIMD unpack for 32-bit values using NEON (fast copy)
90    #[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        // Handle remainder
104        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    /// SIMD prefix sum for 4 u32 values using NEON
113    /// Input:  [a, b, c, d]
114    /// Output: [a, a+b, a+b+c, a+b+c+d]
115    #[inline]
116    #[target_feature(enable = "neon")]
117    unsafe fn prefix_sum_4(v: uint32x4_t) -> uint32x4_t {
118        // Step 1: shift by 1 and add
119        // [a, b, c, d] + [0, a, b, c] = [a, a+b, b+c, c+d]
120        let shifted1 = vextq_u32(vdupq_n_u32(0), v, 3);
121        let sum1 = vaddq_u32(v, shifted1);
122
123        // Step 2: shift by 2 and add
124        // [a, a+b, b+c, c+d] + [0, 0, a, a+b] = [a, a+b, a+b+c, a+b+c+d]
125        let shifted2 = vextq_u32(vdupq_n_u32(0), sum1, 2);
126        vaddq_u32(sum1, shifted2)
127    }
128
129    /// SIMD delta decode: convert deltas to absolute doc IDs
130    /// deltas[i] stores (gap - 1), output[i] = first + sum(gaps[0..i])
131    /// Uses NEON SIMD prefix sum for high throughput
132    #[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            // Load 4 deltas and add 1 (since we store gap-1)
158            let d = vld1q_u32(deltas[base..].as_ptr());
159            let gaps = vaddq_u32(d, ones);
160
161            // Compute prefix sum within the 4 elements
162            let prefix = prefix_sum_4(gaps);
163
164            // Add carry (broadcast last element of previous group)
165            let result = vaddq_u32(prefix, carry);
166
167            // Store result
168            vst1q_u32(output[base + 1..].as_mut_ptr(), result);
169
170            // Update carry: broadcast the last element for next iteration
171            carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
172        }
173
174        // Handle remainder
175        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    /// SIMD add 1 to all values (for TF decoding: stored as tf-1)
184    #[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    /// Fused unpack 8-bit + delta decode using NEON
205    /// Processes 4 values at a time, fusing unpack and prefix sum
206    #[target_feature(enable = "neon")]
207    pub unsafe fn unpack_8bit_delta_decode(
208        input: &[u8],
209        output: &mut [u32],
210        first_value: u32,
211        count: usize,
212    ) {
213        output[0] = first_value;
214        if count <= 1 {
215            return;
216        }
217
218        let ones = vdupq_n_u32(1);
219        let mut carry = vdupq_n_u32(first_value);
220
221        let full_groups = (count - 1) / 4;
222        let remainder = (count - 1) % 4;
223
224        for group in 0..full_groups {
225            let base = group * 4;
226
227            // Load 4 bytes as a u32, then widen u8→u16→u32 via NEON
228            let raw = std::ptr::read_unaligned(input.as_ptr().add(base) as *const u32);
229            let bytes = vreinterpret_u8_u32(vdup_n_u32(raw));
230            let u16s = vmovl_u8(bytes); // 8×u8 → 8×u16 (only low 4 matter)
231            let d = vmovl_u16(vget_low_u16(u16s)); // 4×u16 → 4×u32
232
233            // Add 1 (since we store gap-1)
234            let gaps = vaddq_u32(d, ones);
235
236            // Compute prefix sum within the 4 elements
237            let prefix = prefix_sum_4(gaps);
238
239            // Add carry
240            let result = vaddq_u32(prefix, carry);
241
242            // Store result
243            vst1q_u32(output[base + 1..].as_mut_ptr(), result);
244
245            // Update carry
246            carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
247        }
248
249        // Handle remainder
250        let base = full_groups * 4;
251        let mut scalar_carry = vgetq_lane_u32(carry, 0);
252        for j in 0..remainder {
253            scalar_carry = scalar_carry
254                .wrapping_add(input[base + j] as u32)
255                .wrapping_add(1);
256            output[base + j + 1] = scalar_carry;
257        }
258    }
259
260    /// Fused unpack 16-bit + delta decode using NEON
261    #[target_feature(enable = "neon")]
262    pub unsafe fn unpack_16bit_delta_decode(
263        input: &[u8],
264        output: &mut [u32],
265        first_value: u32,
266        count: usize,
267    ) {
268        output[0] = first_value;
269        if count <= 1 {
270            return;
271        }
272
273        let ones = vdupq_n_u32(1);
274        let mut carry = vdupq_n_u32(first_value);
275
276        let full_groups = (count - 1) / 4;
277        let remainder = (count - 1) % 4;
278
279        for group in 0..full_groups {
280            let base = group * 4;
281            let in_ptr = input.as_ptr().add(base * 2) as *const u16;
282
283            // Load 4 u16 values and widen to u32
284            let vals = vld1_u16(in_ptr);
285            let d = vmovl_u16(vals);
286
287            // Add 1 (since we store gap-1)
288            let gaps = vaddq_u32(d, ones);
289
290            // Compute prefix sum within the 4 elements
291            let prefix = prefix_sum_4(gaps);
292
293            // Add carry
294            let result = vaddq_u32(prefix, carry);
295
296            // Store result
297            vst1q_u32(output[base + 1..].as_mut_ptr(), result);
298
299            // Update carry
300            carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
301        }
302
303        // Handle remainder
304        let base = full_groups * 4;
305        let mut scalar_carry = vgetq_lane_u32(carry, 0);
306        for j in 0..remainder {
307            let idx = (base + j) * 2;
308            let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
309            scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
310            output[base + j + 1] = scalar_carry;
311        }
312    }
313
314    /// NEON Hamming distance: XOR + byte popcount + horizontal sum.
315    /// Processes 16 bytes per iteration (vs 8 for scalar u64 path).
316    #[target_feature(enable = "neon")]
317    pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
318        let len = a.len();
319        let chunks16 = len / 16;
320        let mut total = 0u32;
321
322        // Process 16 bytes at a time, flush u8 accumulators every 31 iters
323        // (vcntq_u8 returns 0-8 per lane; 31 * 8 = 248 ≤ 255, avoiding u8 overflow)
324        let mut i = 0;
325        while i < chunks16 {
326            let batch_end = (i + 31).min(chunks16);
327            let mut acc = vdupq_n_u8(0);
328            for j in i..batch_end {
329                let off = j * 16;
330                let va = vld1q_u8(a.as_ptr().add(off));
331                let vb = vld1q_u8(b.as_ptr().add(off));
332                let popcnt = vcntq_u8(veorq_u8(va, vb));
333                acc = vaddq_u8(acc, popcnt);
334            }
335            // Widen u8 -> u16 -> u32 -> u64 and horizontal sum
336            let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc)));
337            total += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
338            i = batch_end;
339        }
340
341        // Remainder bytes (< 16)
342        let base = chunks16 * 16;
343        for k in base..len {
344            total += (a[k] ^ b[k]).count_ones();
345        }
346
347        total
348    }
349
350    /// Four-row NEON Hamming distance.
351    ///
352    /// The query chunk load and the horizontal reduction are shared across the
353    /// rows, and the four accumulator chains overlap instead of serialising on
354    /// `vcntq_u8`/`vaddq_u8` latency.
355    #[target_feature(enable = "neon")]
356    pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
357        let len = query.len();
358        let chunks16 = len / 16;
359        let mut total = [0u32; 4];
360
361        let mut i = 0;
362        while i < chunks16 {
363            let batch_end = (i + 31).min(chunks16);
364            let mut acc = [vdupq_n_u8(0); 4];
365            for j in i..batch_end {
366                let off = j * 16;
367                let vq = vld1q_u8(query.as_ptr().add(off));
368                for r in 0..4 {
369                    let vr = vld1q_u8(rows[r].as_ptr().add(off));
370                    acc[r] = vaddq_u8(acc[r], vcntq_u8(veorq_u8(vq, vr)));
371                }
372            }
373            for r in 0..4 {
374                let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc[r])));
375                total[r] += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
376            }
377            i = batch_end;
378        }
379
380        // Remainder through the u64 scalar path: a per-byte tail would dominate
381        // for code widths narrower than one vector (e.g. 64-bit fields).
382        let base = chunks16 * 16;
383        if base < len {
384            let tail = &query[base..];
385            for r in 0..4 {
386                total[r] += super::hamming_distance_scalar(tail, &rows[r][base..]);
387            }
388        }
389
390        total
391    }
392
393    /// Check if NEON is available (always true on aarch64)
394    #[inline]
395    pub fn is_available() -> bool {
396        true
397    }
398}
399
400// ============================================================================
401// SSE intrinsics for x86_64 (Intel/AMD)
402// ============================================================================
403
404#[cfg(target_arch = "x86_64")]
405#[allow(unsafe_op_in_unsafe_fn)]
406mod sse {
407    use std::arch::x86_64::*;
408
409    /// SIMD unpack for 8-bit values using SSE
410    #[target_feature(enable = "sse2", enable = "sse4.1")]
411    pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
412        let chunks = count / 16;
413        let remainder = count % 16;
414
415        for chunk in 0..chunks {
416            let base = chunk * 16;
417            let in_ptr = input.as_ptr().add(base);
418
419            let bytes = _mm_loadu_si128(in_ptr as *const __m128i);
420
421            // Zero extend u8 -> u32 using SSE4.1 pmovzx
422            let v0 = _mm_cvtepu8_epi32(bytes);
423            let v1 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 4));
424            let v2 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 8));
425            let v3 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 12));
426
427            let out_ptr = output.as_mut_ptr().add(base);
428            _mm_storeu_si128(out_ptr as *mut __m128i, v0);
429            _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, v1);
430            _mm_storeu_si128(out_ptr.add(8) as *mut __m128i, v2);
431            _mm_storeu_si128(out_ptr.add(12) as *mut __m128i, v3);
432        }
433
434        let base = chunks * 16;
435        for i in 0..remainder {
436            output[base + i] = input[base + i] as u32;
437        }
438    }
439
440    /// SIMD unpack for 16-bit values using SSE
441    #[target_feature(enable = "sse2", enable = "sse4.1")]
442    pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
443        let chunks = count / 8;
444        let remainder = count % 8;
445
446        for chunk in 0..chunks {
447            let base = chunk * 8;
448            let in_ptr = input.as_ptr().add(base * 2);
449
450            let vals = _mm_loadu_si128(in_ptr as *const __m128i);
451            let low = _mm_cvtepu16_epi32(vals);
452            let high = _mm_cvtepu16_epi32(_mm_srli_si128(vals, 8));
453
454            let out_ptr = output.as_mut_ptr().add(base);
455            _mm_storeu_si128(out_ptr as *mut __m128i, low);
456            _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, high);
457        }
458
459        let base = chunks * 8;
460        for i in 0..remainder {
461            let idx = (base + i) * 2;
462            output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
463        }
464    }
465
466    /// SIMD unpack for 32-bit values using SSE (fast copy)
467    #[target_feature(enable = "sse2")]
468    pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
469        let chunks = count / 4;
470        let remainder = count % 4;
471
472        let in_ptr = input.as_ptr() as *const __m128i;
473        let out_ptr = output.as_mut_ptr() as *mut __m128i;
474
475        for chunk in 0..chunks {
476            let vals = _mm_loadu_si128(in_ptr.add(chunk));
477            _mm_storeu_si128(out_ptr.add(chunk), vals);
478        }
479
480        // Handle remainder
481        let base = chunks * 4;
482        for i in 0..remainder {
483            let idx = (base + i) * 4;
484            output[base + i] =
485                u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
486        }
487    }
488
489    /// SIMD prefix sum for 4 u32 values using SSE
490    /// Input:  [a, b, c, d]
491    /// Output: [a, a+b, a+b+c, a+b+c+d]
492    #[inline]
493    #[target_feature(enable = "sse2")]
494    unsafe fn prefix_sum_4(v: __m128i) -> __m128i {
495        // Step 1: shift by 1 element (4 bytes) and add
496        // [a, b, c, d] + [0, a, b, c] = [a, a+b, b+c, c+d]
497        let shifted1 = _mm_slli_si128(v, 4);
498        let sum1 = _mm_add_epi32(v, shifted1);
499
500        // Step 2: shift by 2 elements (8 bytes) and add
501        // [a, a+b, b+c, c+d] + [0, 0, a, a+b] = [a, a+b, a+b+c, a+b+c+d]
502        let shifted2 = _mm_slli_si128(sum1, 8);
503        _mm_add_epi32(sum1, shifted2)
504    }
505
506    /// SIMD delta decode using SSE with true SIMD prefix sum
507    #[target_feature(enable = "sse2", enable = "sse4.1")]
508    pub unsafe fn delta_decode(
509        output: &mut [u32],
510        deltas: &[u32],
511        first_doc_id: u32,
512        count: usize,
513    ) {
514        if count == 0 {
515            return;
516        }
517
518        output[0] = first_doc_id;
519        if count == 1 {
520            return;
521        }
522
523        let ones = _mm_set1_epi32(1);
524        let mut carry = _mm_set1_epi32(first_doc_id as i32);
525
526        let full_groups = (count - 1) / 4;
527        let remainder = (count - 1) % 4;
528
529        for group in 0..full_groups {
530            let base = group * 4;
531
532            // Load 4 deltas and add 1 (since we store gap-1)
533            let d = _mm_loadu_si128(deltas[base..].as_ptr() as *const __m128i);
534            let gaps = _mm_add_epi32(d, ones);
535
536            // Compute prefix sum within the 4 elements
537            let prefix = prefix_sum_4(gaps);
538
539            // Add carry (broadcast last element of previous group)
540            let result = _mm_add_epi32(prefix, carry);
541
542            // Store result
543            _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
544
545            // Update carry: broadcast the last element for next iteration
546            carry = _mm_shuffle_epi32(result, 0xFF); // broadcast lane 3
547        }
548
549        // Handle remainder
550        let base = full_groups * 4;
551        let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
552        for j in 0..remainder {
553            scalar_carry = scalar_carry.wrapping_add(deltas[base + j]).wrapping_add(1);
554            output[base + j + 1] = scalar_carry;
555        }
556    }
557
558    /// SIMD add 1 to all values using SSE
559    #[target_feature(enable = "sse2")]
560    pub unsafe fn add_one(values: &mut [u32], count: usize) {
561        let ones = _mm_set1_epi32(1);
562        let chunks = count / 4;
563        let remainder = count % 4;
564
565        for chunk in 0..chunks {
566            let base = chunk * 4;
567            let ptr = values.as_mut_ptr().add(base) as *mut __m128i;
568            let v = _mm_loadu_si128(ptr);
569            let result = _mm_add_epi32(v, ones);
570            _mm_storeu_si128(ptr, result);
571        }
572
573        let base = chunks * 4;
574        for i in 0..remainder {
575            values[base + i] += 1;
576        }
577    }
578
579    /// Fused unpack 8-bit + delta decode using SSE
580    #[target_feature(enable = "sse2", enable = "sse4.1")]
581    pub unsafe fn unpack_8bit_delta_decode(
582        input: &[u8],
583        output: &mut [u32],
584        first_value: u32,
585        count: usize,
586    ) {
587        output[0] = first_value;
588        if count <= 1 {
589            return;
590        }
591
592        let ones = _mm_set1_epi32(1);
593        let mut carry = _mm_set1_epi32(first_value as i32);
594
595        let full_groups = (count - 1) / 4;
596        let remainder = (count - 1) % 4;
597
598        for group in 0..full_groups {
599            let base = group * 4;
600
601            // Load 4 bytes (unaligned) and zero-extend to u32
602            let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
603                input.as_ptr().add(base) as *const i32
604            ));
605            let d = _mm_cvtepu8_epi32(bytes);
606
607            // Add 1 (since we store gap-1)
608            let gaps = _mm_add_epi32(d, ones);
609
610            // Compute prefix sum within the 4 elements
611            let prefix = prefix_sum_4(gaps);
612
613            // Add carry
614            let result = _mm_add_epi32(prefix, carry);
615
616            // Store result
617            _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
618
619            // Update carry: broadcast the last element
620            carry = _mm_shuffle_epi32(result, 0xFF);
621        }
622
623        // Handle remainder
624        let base = full_groups * 4;
625        let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
626        for j in 0..remainder {
627            scalar_carry = scalar_carry
628                .wrapping_add(input[base + j] as u32)
629                .wrapping_add(1);
630            output[base + j + 1] = scalar_carry;
631        }
632    }
633
634    /// Fused unpack 16-bit + delta decode using SSE
635    #[target_feature(enable = "sse2", enable = "sse4.1")]
636    pub unsafe fn unpack_16bit_delta_decode(
637        input: &[u8],
638        output: &mut [u32],
639        first_value: u32,
640        count: usize,
641    ) {
642        output[0] = first_value;
643        if count <= 1 {
644            return;
645        }
646
647        let ones = _mm_set1_epi32(1);
648        let mut carry = _mm_set1_epi32(first_value as i32);
649
650        let full_groups = (count - 1) / 4;
651        let remainder = (count - 1) % 4;
652
653        for group in 0..full_groups {
654            let base = group * 4;
655            let in_ptr = input.as_ptr().add(base * 2);
656
657            // Load 8 bytes (4 u16 values, unaligned) and zero-extend to u32
658            let vals = _mm_loadl_epi64(in_ptr as *const __m128i); // loadl_epi64 supports unaligned
659            let d = _mm_cvtepu16_epi32(vals);
660
661            // Add 1 (since we store gap-1)
662            let gaps = _mm_add_epi32(d, ones);
663
664            // Compute prefix sum within the 4 elements
665            let prefix = prefix_sum_4(gaps);
666
667            // Add carry
668            let result = _mm_add_epi32(prefix, carry);
669
670            // Store result
671            _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
672
673            // Update carry: broadcast the last element
674            carry = _mm_shuffle_epi32(result, 0xFF);
675        }
676
677        // Handle remainder
678        let base = full_groups * 4;
679        let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
680        for j in 0..remainder {
681            let idx = (base + j) * 2;
682            let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
683            scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
684            output[base + j + 1] = scalar_carry;
685        }
686    }
687
688    /// Check if SSE4.1 is available at runtime
689    #[inline]
690    pub fn is_available() -> bool {
691        is_x86_feature_detected!("sse4.1")
692    }
693}
694
695// ============================================================================
696// AVX2 intrinsics for x86_64 (Intel/AMD with 256-bit registers)
697// ============================================================================
698
699#[cfg(target_arch = "x86_64")]
700#[allow(unsafe_op_in_unsafe_fn)]
701mod avx2 {
702    use std::arch::x86_64::*;
703
704    /// AVX2 unpack for 8-bit values (processes 32 bytes at a time)
705    #[target_feature(enable = "avx2")]
706    pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
707        let chunks = count / 32;
708        let remainder = count % 32;
709
710        for chunk in 0..chunks {
711            let base = chunk * 32;
712            let in_ptr = input.as_ptr().add(base);
713
714            // Load 32 bytes (two 128-bit loads, then combine)
715            let bytes_lo = _mm_loadu_si128(in_ptr as *const __m128i);
716            let bytes_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
717
718            // Zero extend first 16 bytes: u8 -> u32
719            let v0 = _mm256_cvtepu8_epi32(bytes_lo);
720            let v1 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_lo, 8));
721            let v2 = _mm256_cvtepu8_epi32(bytes_hi);
722            let v3 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_hi, 8));
723
724            let out_ptr = output.as_mut_ptr().add(base);
725            _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
726            _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
727            _mm256_storeu_si256(out_ptr.add(16) as *mut __m256i, v2);
728            _mm256_storeu_si256(out_ptr.add(24) as *mut __m256i, v3);
729        }
730
731        // Handle remainder with SSE
732        let base = chunks * 32;
733        for i in 0..remainder {
734            output[base + i] = input[base + i] as u32;
735        }
736    }
737
738    /// AVX2 unpack for 16-bit values (processes 16 values at a time)
739    #[target_feature(enable = "avx2")]
740    pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
741        let chunks = count / 16;
742        let remainder = count % 16;
743
744        for chunk in 0..chunks {
745            let base = chunk * 16;
746            let in_ptr = input.as_ptr().add(base * 2);
747
748            // Load 32 bytes (16 u16 values)
749            let vals_lo = _mm_loadu_si128(in_ptr as *const __m128i);
750            let vals_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
751
752            // Zero extend u16 -> u32
753            let v0 = _mm256_cvtepu16_epi32(vals_lo);
754            let v1 = _mm256_cvtepu16_epi32(vals_hi);
755
756            let out_ptr = output.as_mut_ptr().add(base);
757            _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
758            _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
759        }
760
761        // Handle remainder
762        let base = chunks * 16;
763        for i in 0..remainder {
764            let idx = (base + i) * 2;
765            output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
766        }
767    }
768
769    /// AVX2 unpack for 32-bit values (fast copy, 8 values at a time)
770    #[target_feature(enable = "avx2")]
771    pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
772        let chunks = count / 8;
773        let remainder = count % 8;
774
775        let in_ptr = input.as_ptr() as *const __m256i;
776        let out_ptr = output.as_mut_ptr() as *mut __m256i;
777
778        for chunk in 0..chunks {
779            let vals = _mm256_loadu_si256(in_ptr.add(chunk));
780            _mm256_storeu_si256(out_ptr.add(chunk), vals);
781        }
782
783        // Handle remainder
784        let base = chunks * 8;
785        for i in 0..remainder {
786            let idx = (base + i) * 4;
787            output[base + i] =
788                u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
789        }
790    }
791
792    /// AVX2 add 1 to all values (8 values at a time)
793    #[target_feature(enable = "avx2")]
794    pub unsafe fn add_one(values: &mut [u32], count: usize) {
795        let ones = _mm256_set1_epi32(1);
796        let chunks = count / 8;
797        let remainder = count % 8;
798
799        for chunk in 0..chunks {
800            let base = chunk * 8;
801            let ptr = values.as_mut_ptr().add(base) as *mut __m256i;
802            let v = _mm256_loadu_si256(ptr);
803            let result = _mm256_add_epi32(v, ones);
804            _mm256_storeu_si256(ptr, result);
805        }
806
807        let base = chunks * 8;
808        for i in 0..remainder {
809            values[base + i] += 1;
810        }
811    }
812
813    /// AVX2 prefix sum for 8 u32 values (Hillis-Steele)
814    /// Input:  [a, b, c, d, e, f, g, h]
815    /// Output: [a, a+b, a+b+c, ..., a+b+c+d+e+f+g+h]
816    #[inline]
817    #[target_feature(enable = "avx2")]
818    unsafe fn prefix_sum_8(v: __m256i) -> __m256i {
819        // Step 1: intra-lane shift by 1 element (4 bytes) and add
820        let s1 = _mm256_slli_si256(v, 4);
821        let r1 = _mm256_add_epi32(v, s1);
822
823        // Step 2: intra-lane shift by 2 elements (8 bytes) and add
824        let s2 = _mm256_slli_si256(r1, 8);
825        let r2 = _mm256_add_epi32(r1, s2);
826
827        // Step 3: propagate lower lane sum to upper lane
828        // Broadcast element 3 (lower lane sum) within each lane
829        let lo_sum = _mm256_shuffle_epi32(r2, 0xFF);
830        // Duplicate lane 0 to both lanes
831        let carry = _mm256_permute2x128_si256(lo_sum, lo_sum, 0x00);
832        // Zero carry for lower lane, keep for upper
833        let carry_hi = _mm256_blend_epi32::<0xF0>(_mm256_setzero_si256(), carry);
834        _mm256_add_epi32(r2, carry_hi)
835    }
836
837    /// AVX2 fused unpack 8-bit + delta decode (processes 8 values at a time)
838    #[target_feature(enable = "avx2")]
839    pub unsafe fn unpack_8bit_delta_decode(
840        input: &[u8],
841        output: &mut [u32],
842        first_value: u32,
843        count: usize,
844    ) {
845        output[0] = first_value;
846        if count <= 1 {
847            return;
848        }
849
850        let ones = _mm256_set1_epi32(1);
851        let mut carry = _mm256_set1_epi32(first_value as i32);
852        let broadcast_idx = _mm256_set1_epi32(7);
853
854        let full_groups = (count - 1) / 8;
855        let remainder = (count - 1) % 8;
856
857        for group in 0..full_groups {
858            let base = group * 8;
859
860            // Load 8 bytes and zero-extend to 8×u32
861            let bytes = _mm_loadl_epi64(input.as_ptr().add(base) as *const __m128i);
862            let d = _mm256_cvtepu8_epi32(bytes);
863
864            // Add 1 (since we store gap-1)
865            let gaps = _mm256_add_epi32(d, ones);
866
867            // Compute prefix sum within 8 elements
868            let prefix = prefix_sum_8(gaps);
869
870            // Add carry from previous group
871            let result = _mm256_add_epi32(prefix, carry);
872
873            // Store 8 results
874            _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
875
876            // Update carry: broadcast element 7 to all positions
877            carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
878        }
879
880        // Handle remainder with scalar
881        let base = full_groups * 8;
882        let mut scalar_carry = _mm256_extract_epi32::<0>(carry) as u32;
883        for j in 0..remainder {
884            scalar_carry = scalar_carry
885                .wrapping_add(input[base + j] as u32)
886                .wrapping_add(1);
887            output[base + j + 1] = scalar_carry;
888        }
889    }
890
891    /// AVX2 fused unpack 16-bit + delta decode (processes 8 values at a time)
892    #[target_feature(enable = "avx2")]
893    pub unsafe fn unpack_16bit_delta_decode(
894        input: &[u8],
895        output: &mut [u32],
896        first_value: u32,
897        count: usize,
898    ) {
899        output[0] = first_value;
900        if count <= 1 {
901            return;
902        }
903
904        let ones = _mm256_set1_epi32(1);
905        let mut carry = _mm256_set1_epi32(first_value as i32);
906        let broadcast_idx = _mm256_set1_epi32(7);
907
908        let full_groups = (count - 1) / 8;
909        let remainder = (count - 1) % 8;
910
911        for group in 0..full_groups {
912            let base = group * 8;
913            let in_ptr = input.as_ptr().add(base * 2);
914
915            // Load 16 bytes (8 u16 values) and zero-extend to 8×u32
916            let vals = _mm_loadu_si128(in_ptr as *const __m128i);
917            let d = _mm256_cvtepu16_epi32(vals);
918
919            // Add 1 (since we store gap-1)
920            let gaps = _mm256_add_epi32(d, ones);
921
922            // Compute prefix sum within 8 elements
923            let prefix = prefix_sum_8(gaps);
924
925            // Add carry from previous group
926            let result = _mm256_add_epi32(prefix, carry);
927
928            // Store 8 results
929            _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
930
931            // Update carry: broadcast element 7 to all positions
932            carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
933        }
934
935        // Handle remainder with scalar
936        let base = full_groups * 8;
937        let mut scalar_carry = _mm256_extract_epi32::<0>(carry) as u32;
938        for j in 0..remainder {
939            let idx = (base + j) * 2;
940            let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
941            scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
942            output[base + j + 1] = scalar_carry;
943        }
944    }
945
946    /// AVX2 Hamming distance using VPSHUFB-based popcount (Muła algorithm).
947    /// Processes 32 bytes per iteration with a nibble lookup table.
948    #[target_feature(enable = "avx2")]
949    pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
950        let len = a.len();
951        let chunks32 = len / 32;
952        let low_mask = _mm256_set1_epi8(0x0f);
953        // Nibble popcount lookup table: popcount(0..15)
954        let lookup = _mm256_setr_epi8(
955            0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4, 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2,
956            3, 3, 4,
957        );
958        let mut total = 0u64;
959
960        let mut i = 0;
961        while i < chunks32 {
962            // Accumulate in u8 lanes, flush every 31 iters to avoid overflow
963            // (nibble popcount gives 0-8 per lane; 31 * 8 = 248 ≤ 255)
964            let batch_end = (i + 31).min(chunks32);
965            let mut acc = _mm256_setzero_si256();
966            for j in i..batch_end {
967                let off = j * 32;
968                let va = _mm256_loadu_si256(a.as_ptr().add(off) as *const __m256i);
969                let vb = _mm256_loadu_si256(b.as_ptr().add(off) as *const __m256i);
970                let xored = _mm256_xor_si256(va, vb);
971                // VPSHUFB popcount: count bits per byte via nibble lookup
972                let lo = _mm256_and_si256(xored, low_mask);
973                let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
974                let popcnt = _mm256_add_epi8(
975                    _mm256_shuffle_epi8(lookup, lo),
976                    _mm256_shuffle_epi8(lookup, hi),
977                );
978                acc = _mm256_add_epi8(acc, popcnt);
979            }
980            // Horizontal sum: u8 -> u64 via SAD against zero
981            let sad = _mm256_sad_epu8(acc, _mm256_setzero_si256());
982            total += _mm256_extract_epi64(sad, 0) as u64
983                + _mm256_extract_epi64(sad, 1) as u64
984                + _mm256_extract_epi64(sad, 2) as u64
985                + _mm256_extract_epi64(sad, 3) as u64;
986            i = batch_end;
987        }
988
989        // Remainder bytes (< 32)
990        let base = chunks32 * 32;
991        for k in base..len {
992            total += (a[k] ^ b[k]).count_ones() as u64;
993        }
994
995        total as u32
996    }
997
998    /// Four-row AVX2 Hamming distance.
999    ///
1000    /// The query chunk load, the nibble lookup table and the horizontal
1001    /// reduction are shared across the rows, and the four accumulator chains
1002    /// overlap instead of serialising on popcount latency.
1003    #[target_feature(enable = "avx2")]
1004    pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
1005        let len = query.len();
1006        let chunks32 = len / 32;
1007        let low_mask = _mm256_set1_epi8(0x0f);
1008        let lookup = _mm256_setr_epi8(
1009            0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4, 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2,
1010            3, 3, 4,
1011        );
1012        let mut total = [0u64; 4];
1013
1014        let mut i = 0;
1015        while i < chunks32 {
1016            let batch_end = (i + 31).min(chunks32);
1017            let mut acc = [_mm256_setzero_si256(); 4];
1018            for j in i..batch_end {
1019                let off = j * 32;
1020                let vq = _mm256_loadu_si256(query.as_ptr().add(off) as *const __m256i);
1021                for r in 0..4 {
1022                    let vr = _mm256_loadu_si256(rows[r].as_ptr().add(off) as *const __m256i);
1023                    let xored = _mm256_xor_si256(vq, vr);
1024                    let lo = _mm256_and_si256(xored, low_mask);
1025                    let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
1026                    acc[r] = _mm256_add_epi8(
1027                        acc[r],
1028                        _mm256_add_epi8(
1029                            _mm256_shuffle_epi8(lookup, lo),
1030                            _mm256_shuffle_epi8(lookup, hi),
1031                        ),
1032                    );
1033                }
1034            }
1035            for r in 0..4 {
1036                let sad = _mm256_sad_epu8(acc[r], _mm256_setzero_si256());
1037                total[r] += _mm256_extract_epi64(sad, 0) as u64
1038                    + _mm256_extract_epi64(sad, 1) as u64
1039                    + _mm256_extract_epi64(sad, 2) as u64
1040                    + _mm256_extract_epi64(sad, 3) as u64;
1041            }
1042            i = batch_end;
1043        }
1044
1045        // Remainder through the u64 scalar path: a per-byte tail would dominate
1046        // for code widths narrower than one vector (e.g. 64-bit fields).
1047        let base = chunks32 * 32;
1048        if base < len {
1049            let tail = &query[base..];
1050            for r in 0..4 {
1051                total[r] += u64::from(super::hamming_distance_scalar(tail, &rows[r][base..]));
1052            }
1053        }
1054
1055        [
1056            total[0] as u32,
1057            total[1] as u32,
1058            total[2] as u32,
1059            total[3] as u32,
1060        ]
1061    }
1062
1063    /// Check if AVX2 is available at runtime
1064    #[inline]
1065    pub fn is_available() -> bool {
1066        is_x86_feature_detected!("avx2")
1067    }
1068}
1069
1070// ============================================================================
1071// Scalar fallback implementations
1072// ============================================================================
1073
1074#[allow(dead_code)]
1075mod scalar {
1076    /// Scalar unpack for 8-bit values
1077    #[inline]
1078    pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1079        for i in 0..count {
1080            output[i] = input[i] as u32;
1081        }
1082    }
1083
1084    /// Scalar unpack for 16-bit values
1085    #[inline]
1086    pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1087        for (i, out) in output.iter_mut().enumerate().take(count) {
1088            let idx = i * 2;
1089            *out = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
1090        }
1091    }
1092
1093    /// Scalar unpack for 32-bit values
1094    #[inline]
1095    pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1096        for (i, out) in output.iter_mut().enumerate().take(count) {
1097            let idx = i * 4;
1098            *out = u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
1099        }
1100    }
1101
1102    /// Scalar delta decode
1103    #[inline]
1104    pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_doc_id: u32, count: usize) {
1105        if count == 0 {
1106            return;
1107        }
1108
1109        output[0] = first_doc_id;
1110        let mut carry = first_doc_id;
1111
1112        for i in 0..count - 1 {
1113            carry = carry.wrapping_add(deltas[i]).wrapping_add(1);
1114            output[i + 1] = carry;
1115        }
1116    }
1117
1118    /// Scalar add 1 to all values
1119    #[inline]
1120    pub fn add_one(values: &mut [u32], count: usize) {
1121        for val in values.iter_mut().take(count) {
1122            *val += 1;
1123        }
1124    }
1125}
1126
1127// ============================================================================
1128// Public dispatch functions that select SIMD or scalar at runtime
1129// ============================================================================
1130
1131/// Unpack 8-bit packed values to u32 with SIMD acceleration
1132#[inline]
1133pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1134    #[cfg(target_arch = "aarch64")]
1135    {
1136        if neon::is_available() {
1137            unsafe {
1138                neon::unpack_8bit(input, output, count);
1139            }
1140            return;
1141        }
1142    }
1143
1144    #[cfg(target_arch = "x86_64")]
1145    {
1146        // Prefer AVX2 (256-bit) over SSE (128-bit) when available
1147        if avx2::is_available() {
1148            unsafe {
1149                avx2::unpack_8bit(input, output, count);
1150            }
1151            return;
1152        }
1153        if sse::is_available() {
1154            unsafe {
1155                sse::unpack_8bit(input, output, count);
1156            }
1157            return;
1158        }
1159    }
1160
1161    scalar::unpack_8bit(input, output, count);
1162}
1163
1164/// Unpack 16-bit packed values to u32 with SIMD acceleration
1165#[inline]
1166pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1167    #[cfg(target_arch = "aarch64")]
1168    {
1169        if neon::is_available() {
1170            unsafe {
1171                neon::unpack_16bit(input, output, count);
1172            }
1173            return;
1174        }
1175    }
1176
1177    #[cfg(target_arch = "x86_64")]
1178    {
1179        // Prefer AVX2 (256-bit) over SSE (128-bit) when available
1180        if avx2::is_available() {
1181            unsafe {
1182                avx2::unpack_16bit(input, output, count);
1183            }
1184            return;
1185        }
1186        if sse::is_available() {
1187            unsafe {
1188                sse::unpack_16bit(input, output, count);
1189            }
1190            return;
1191        }
1192    }
1193
1194    scalar::unpack_16bit(input, output, count);
1195}
1196
1197/// Unpack 32-bit packed values to u32 with SIMD acceleration
1198#[inline]
1199pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1200    #[cfg(target_arch = "aarch64")]
1201    {
1202        if neon::is_available() {
1203            unsafe {
1204                neon::unpack_32bit(input, output, count);
1205            }
1206            return;
1207        }
1208    }
1209
1210    #[cfg(target_arch = "x86_64")]
1211    {
1212        // Prefer AVX2 (256-bit) over SSE (128-bit) when available
1213        if avx2::is_available() {
1214            unsafe {
1215                avx2::unpack_32bit(input, output, count);
1216            }
1217            return;
1218        }
1219        if sse::is_available() {
1220            unsafe {
1221                sse::unpack_32bit(input, output, count);
1222            }
1223            return;
1224        }
1225    }
1226
1227    scalar::unpack_32bit(input, output, count);
1228}
1229
1230/// Delta decode with SIMD acceleration
1231///
1232/// Converts delta-encoded values to absolute values.
1233/// Input: `deltas[i] = value[i + 1] - value[i] - 1` (gap minus one)
1234/// Output: absolute values starting from first_value
1235#[inline]
1236pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_value: u32, count: usize) {
1237    #[cfg(target_arch = "aarch64")]
1238    {
1239        if neon::is_available() {
1240            unsafe {
1241                neon::delta_decode(output, deltas, first_value, count);
1242            }
1243            return;
1244        }
1245    }
1246
1247    #[cfg(target_arch = "x86_64")]
1248    {
1249        if sse::is_available() {
1250            unsafe {
1251                sse::delta_decode(output, deltas, first_value, count);
1252            }
1253            return;
1254        }
1255    }
1256
1257    scalar::delta_decode(output, deltas, first_value, count);
1258}
1259
1260/// Add 1 to all values with SIMD acceleration
1261///
1262/// Used for TF decoding where values are stored as (tf - 1)
1263#[inline]
1264pub fn add_one(values: &mut [u32], count: usize) {
1265    #[cfg(target_arch = "aarch64")]
1266    {
1267        if neon::is_available() {
1268            unsafe {
1269                neon::add_one(values, count);
1270            }
1271            return;
1272        }
1273    }
1274
1275    #[cfg(target_arch = "x86_64")]
1276    {
1277        // Prefer AVX2 (256-bit) over SSE (128-bit) when available
1278        if avx2::is_available() {
1279            unsafe {
1280                avx2::add_one(values, count);
1281            }
1282            return;
1283        }
1284        if sse::is_available() {
1285            unsafe {
1286                sse::add_one(values, count);
1287            }
1288            return;
1289        }
1290    }
1291
1292    scalar::add_one(values, count);
1293}
1294
1295/// Compute the number of bits needed to represent a value
1296#[inline]
1297pub fn bits_needed(val: u32) -> u8 {
1298    if val == 0 {
1299        0
1300    } else {
1301        32 - val.leading_zeros() as u8
1302    }
1303}
1304
1305// ============================================================================
1306// Rounded bitpacking for truly vectorized encoding/decoding
1307// ============================================================================
1308//
1309// Instead of using arbitrary bit widths (1-32), we round up to SIMD-friendly
1310// widths: 0, 8, 16, or 32 bits. This trades ~10-20% more space for much faster
1311// decoding since we can use direct SIMD widening instructions (pmovzx) without
1312// any bit-shifting or masking.
1313//
1314// Bit width mapping:
1315//   0      -> 0  (all zeros)
1316//   1-8    -> 8  (u8)
1317//   9-16   -> 16 (u16)
1318//   17-32  -> 32 (u32)
1319
1320/// Rounded bit width type for SIMD-friendly encoding
1321#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1322#[repr(u8)]
1323pub enum RoundedBitWidth {
1324    Zero = 0,
1325    Bits8 = 8,
1326    Bits16 = 16,
1327    Bits32 = 32,
1328}
1329
1330impl RoundedBitWidth {
1331    /// Round an exact bit width to the nearest SIMD-friendly width
1332    #[inline]
1333    pub fn from_exact(bits: u8) -> Self {
1334        match bits {
1335            0 => RoundedBitWidth::Zero,
1336            1..=8 => RoundedBitWidth::Bits8,
1337            9..=16 => RoundedBitWidth::Bits16,
1338            _ => RoundedBitWidth::Bits32,
1339        }
1340    }
1341
1342    /// Convert from stored u8 value (must be 0, 8, 16, or 32)
1343    #[inline]
1344    pub fn from_u8(bits: u8) -> Self {
1345        match bits {
1346            0 => RoundedBitWidth::Zero,
1347            8 => RoundedBitWidth::Bits8,
1348            16 => RoundedBitWidth::Bits16,
1349            32 => RoundedBitWidth::Bits32,
1350            _ => RoundedBitWidth::Bits32, // Fallback for invalid values
1351        }
1352    }
1353
1354    /// Get the byte size per value
1355    #[inline]
1356    pub fn bytes_per_value(self) -> usize {
1357        match self {
1358            RoundedBitWidth::Zero => 0,
1359            RoundedBitWidth::Bits8 => 1,
1360            RoundedBitWidth::Bits16 => 2,
1361            RoundedBitWidth::Bits32 => 4,
1362        }
1363    }
1364
1365    /// Get the raw bit width value
1366    #[inline]
1367    pub fn as_u8(self) -> u8 {
1368        self as u8
1369    }
1370}
1371
1372/// Round a bit width to the nearest SIMD-friendly width (0, 8, 16, or 32)
1373#[inline]
1374pub fn round_bit_width(bits: u8) -> u8 {
1375    RoundedBitWidth::from_exact(bits).as_u8()
1376}
1377
1378/// Pack values using rounded bit width (SIMD-friendly)
1379///
1380/// This is much simpler than arbitrary bitpacking since values are byte-aligned.
1381/// Returns the number of bytes written.
1382#[inline]
1383pub fn pack_rounded(values: &[u32], bit_width: RoundedBitWidth, output: &mut [u8]) -> usize {
1384    let count = values.len();
1385    match bit_width {
1386        RoundedBitWidth::Zero => 0,
1387        RoundedBitWidth::Bits8 => {
1388            for (i, &v) in values.iter().enumerate() {
1389                output[i] = v as u8;
1390            }
1391            count
1392        }
1393        RoundedBitWidth::Bits16 => {
1394            for (i, &v) in values.iter().enumerate() {
1395                let bytes = (v as u16).to_le_bytes();
1396                output[i * 2] = bytes[0];
1397                output[i * 2 + 1] = bytes[1];
1398            }
1399            count * 2
1400        }
1401        RoundedBitWidth::Bits32 => {
1402            for (i, &v) in values.iter().enumerate() {
1403                let bytes = v.to_le_bytes();
1404                output[i * 4] = bytes[0];
1405                output[i * 4 + 1] = bytes[1];
1406                output[i * 4 + 2] = bytes[2];
1407                output[i * 4 + 3] = bytes[3];
1408            }
1409            count * 4
1410        }
1411    }
1412}
1413
1414/// Unpack values using rounded bit width with SIMD acceleration
1415///
1416/// This is the fast path - no bit manipulation needed, just widening.
1417#[inline]
1418pub fn unpack_rounded(input: &[u8], bit_width: RoundedBitWidth, output: &mut [u32], count: usize) {
1419    match bit_width {
1420        RoundedBitWidth::Zero => {
1421            for out in output.iter_mut().take(count) {
1422                *out = 0;
1423            }
1424        }
1425        RoundedBitWidth::Bits8 => unpack_8bit(input, output, count),
1426        RoundedBitWidth::Bits16 => unpack_16bit(input, output, count),
1427        RoundedBitWidth::Bits32 => unpack_32bit(input, output, count),
1428    }
1429}
1430
1431/// Fused unpack + delta decode using rounded bit width
1432///
1433/// Combines unpacking and prefix sum in a single pass for better cache utilization.
1434#[inline]
1435pub fn unpack_rounded_delta_decode(
1436    input: &[u8],
1437    bit_width: RoundedBitWidth,
1438    output: &mut [u32],
1439    first_value: u32,
1440    count: usize,
1441) {
1442    match bit_width {
1443        RoundedBitWidth::Zero => {
1444            // All deltas are 0, meaning gaps of 1
1445            let mut val = first_value;
1446            for out in output.iter_mut().take(count) {
1447                *out = val;
1448                val = val.wrapping_add(1);
1449            }
1450        }
1451        RoundedBitWidth::Bits8 => unpack_8bit_delta_decode(input, output, first_value, count),
1452        RoundedBitWidth::Bits16 => unpack_16bit_delta_decode(input, output, first_value, count),
1453        RoundedBitWidth::Bits32 => {
1454            // Unpack count-1 deltas from input, then prefix sum to absolute values
1455            if count > 0 {
1456                output[0] = first_value;
1457                let mut carry = first_value;
1458                for i in 0..count - 1 {
1459                    let idx = i * 4;
1460                    let delta = u32::from_le_bytes([
1461                        input[idx],
1462                        input[idx + 1],
1463                        input[idx + 2],
1464                        input[idx + 3],
1465                    ]);
1466                    carry = carry.wrapping_add(delta).wrapping_add(1);
1467                    output[i + 1] = carry;
1468                }
1469            }
1470        }
1471    }
1472}
1473
1474// ============================================================================
1475// Fused operations for better cache utilization
1476// ============================================================================
1477
1478/// Fused unpack 8-bit + delta decode in a single pass
1479///
1480/// This avoids writing the intermediate unpacked values to memory,
1481/// improving cache utilization for large blocks.
1482#[inline]
1483pub fn unpack_8bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1484    if count == 0 {
1485        return;
1486    }
1487
1488    output[0] = first_value;
1489    if count == 1 {
1490        return;
1491    }
1492
1493    #[cfg(target_arch = "aarch64")]
1494    {
1495        if neon::is_available() {
1496            unsafe {
1497                neon::unpack_8bit_delta_decode(input, output, first_value, count);
1498            }
1499            return;
1500        }
1501    }
1502
1503    #[cfg(target_arch = "x86_64")]
1504    {
1505        if avx2::is_available() {
1506            unsafe {
1507                avx2::unpack_8bit_delta_decode(input, output, first_value, count);
1508            }
1509            return;
1510        }
1511        if sse::is_available() {
1512            unsafe {
1513                sse::unpack_8bit_delta_decode(input, output, first_value, count);
1514            }
1515            return;
1516        }
1517    }
1518
1519    // Scalar fallback
1520    let mut carry = first_value;
1521    for i in 0..count - 1 {
1522        carry = carry.wrapping_add(input[i] as u32).wrapping_add(1);
1523        output[i + 1] = carry;
1524    }
1525}
1526
1527/// Fused unpack 16-bit + delta decode in a single pass
1528#[inline]
1529pub fn unpack_16bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1530    if count == 0 {
1531        return;
1532    }
1533
1534    output[0] = first_value;
1535    if count == 1 {
1536        return;
1537    }
1538
1539    #[cfg(target_arch = "aarch64")]
1540    {
1541        if neon::is_available() {
1542            unsafe {
1543                neon::unpack_16bit_delta_decode(input, output, first_value, count);
1544            }
1545            return;
1546        }
1547    }
1548
1549    #[cfg(target_arch = "x86_64")]
1550    {
1551        if avx2::is_available() {
1552            unsafe {
1553                avx2::unpack_16bit_delta_decode(input, output, first_value, count);
1554            }
1555            return;
1556        }
1557        if sse::is_available() {
1558            unsafe {
1559                sse::unpack_16bit_delta_decode(input, output, first_value, count);
1560            }
1561            return;
1562        }
1563    }
1564
1565    // Scalar fallback
1566    let mut carry = first_value;
1567    for i in 0..count - 1 {
1568        let idx = i * 2;
1569        let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
1570        carry = carry.wrapping_add(delta).wrapping_add(1);
1571        output[i + 1] = carry;
1572    }
1573}
1574
1575/// Fused unpack + delta decode for arbitrary bit widths
1576///
1577/// Combines unpacking and prefix sum in a single pass, avoiding intermediate buffer.
1578/// Uses SIMD-accelerated paths for 8/16-bit widths, scalar for others.
1579#[inline]
1580pub fn unpack_delta_decode(
1581    input: &[u8],
1582    bit_width: u8,
1583    output: &mut [u32],
1584    first_value: u32,
1585    count: usize,
1586) {
1587    if count == 0 {
1588        return;
1589    }
1590
1591    output[0] = first_value;
1592    if count == 1 {
1593        return;
1594    }
1595
1596    // Fast paths for SIMD-friendly bit widths
1597    match bit_width {
1598        0 => {
1599            // All zeros = consecutive doc IDs (gap of 1)
1600            let mut val = first_value;
1601            for item in output.iter_mut().take(count).skip(1) {
1602                val = val.wrapping_add(1);
1603                *item = val;
1604            }
1605        }
1606        8 => unpack_8bit_delta_decode(input, output, first_value, count),
1607        16 => unpack_16bit_delta_decode(input, output, first_value, count),
1608        32 => {
1609            // 32-bit: unpack inline and delta decode
1610            let mut carry = first_value;
1611            for i in 0..count - 1 {
1612                let idx = i * 4;
1613                let delta = u32::from_le_bytes([
1614                    input[idx],
1615                    input[idx + 1],
1616                    input[idx + 2],
1617                    input[idx + 3],
1618                ]);
1619                carry = carry.wrapping_add(delta).wrapping_add(1);
1620                output[i + 1] = carry;
1621            }
1622        }
1623        _ => {
1624            // Generic bit width: fused unpack + delta decode
1625            let mask = (1u64 << bit_width) - 1;
1626            let bit_width_usize = bit_width as usize;
1627            let mut bit_pos = 0usize;
1628            let input_ptr = input.as_ptr();
1629            let mut carry = first_value;
1630
1631            for i in 0..count - 1 {
1632                let byte_idx = bit_pos >> 3;
1633                let bit_offset = bit_pos & 7;
1634
1635                // SAFETY: Caller guarantees input has enough data
1636                let word = unsafe { (input_ptr.add(byte_idx) as *const u64).read_unaligned() };
1637                let delta = ((word >> bit_offset) & mask) as u32;
1638
1639                carry = carry.wrapping_add(delta).wrapping_add(1);
1640                output[i + 1] = carry;
1641                bit_pos += bit_width_usize;
1642            }
1643        }
1644    }
1645}
1646
1647// ============================================================================
1648// Sparse Vector SIMD Functions
1649// ============================================================================
1650
1651/// Dequantize UInt8 weights to f32 with SIMD acceleration
1652///
1653/// Computes `output[i] = input[i] as f32 * scale + min_val`.
1654#[inline]
1655pub fn dequantize_uint8(input: &[u8], output: &mut [f32], scale: f32, min_val: f32, count: usize) {
1656    #[cfg(target_arch = "aarch64")]
1657    {
1658        if neon::is_available() {
1659            unsafe {
1660                dequantize_uint8_neon(input, output, scale, min_val, count);
1661            }
1662            return;
1663        }
1664    }
1665
1666    #[cfg(target_arch = "x86_64")]
1667    {
1668        if sse::is_available() {
1669            unsafe {
1670                dequantize_uint8_sse(input, output, scale, min_val, count);
1671            }
1672            return;
1673        }
1674    }
1675
1676    // Scalar fallback
1677    for i in 0..count {
1678        output[i] = input[i] as f32 * scale + min_val;
1679    }
1680}
1681
1682#[cfg(target_arch = "aarch64")]
1683#[target_feature(enable = "neon")]
1684#[allow(unsafe_op_in_unsafe_fn)]
1685unsafe fn dequantize_uint8_neon(
1686    input: &[u8],
1687    output: &mut [f32],
1688    scale: f32,
1689    min_val: f32,
1690    count: usize,
1691) {
1692    use std::arch::aarch64::*;
1693
1694    let scale_v = vdupq_n_f32(scale);
1695    let min_v = vdupq_n_f32(min_val);
1696
1697    let chunks = count / 16;
1698    let remainder = count % 16;
1699
1700    for chunk in 0..chunks {
1701        let base = chunk * 16;
1702        let in_ptr = input.as_ptr().add(base);
1703
1704        // Load 16 bytes
1705        let bytes = vld1q_u8(in_ptr);
1706
1707        // Widen u8 -> u16 -> u32 -> f32
1708        let low8 = vget_low_u8(bytes);
1709        let high8 = vget_high_u8(bytes);
1710
1711        let low16 = vmovl_u8(low8);
1712        let high16 = vmovl_u8(high8);
1713
1714        // Process 4 values at a time
1715        let u32_0 = vmovl_u16(vget_low_u16(low16));
1716        let u32_1 = vmovl_u16(vget_high_u16(low16));
1717        let u32_2 = vmovl_u16(vget_low_u16(high16));
1718        let u32_3 = vmovl_u16(vget_high_u16(high16));
1719
1720        // Convert to f32 and apply scale + min_val
1721        let f32_0 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_0), scale_v);
1722        let f32_1 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_1), scale_v);
1723        let f32_2 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_2), scale_v);
1724        let f32_3 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_3), scale_v);
1725
1726        let out_ptr = output.as_mut_ptr().add(base);
1727        vst1q_f32(out_ptr, f32_0);
1728        vst1q_f32(out_ptr.add(4), f32_1);
1729        vst1q_f32(out_ptr.add(8), f32_2);
1730        vst1q_f32(out_ptr.add(12), f32_3);
1731    }
1732
1733    // Handle remainder
1734    let base = chunks * 16;
1735    for i in 0..remainder {
1736        output[base + i] = input[base + i] as f32 * scale + min_val;
1737    }
1738}
1739
1740#[cfg(target_arch = "x86_64")]
1741#[target_feature(enable = "sse2", enable = "sse4.1")]
1742#[allow(unsafe_op_in_unsafe_fn)]
1743unsafe fn dequantize_uint8_sse(
1744    input: &[u8],
1745    output: &mut [f32],
1746    scale: f32,
1747    min_val: f32,
1748    count: usize,
1749) {
1750    use std::arch::x86_64::*;
1751
1752    let scale_v = _mm_set1_ps(scale);
1753    let min_v = _mm_set1_ps(min_val);
1754
1755    let chunks = count / 4;
1756    let remainder = count % 4;
1757
1758    for chunk in 0..chunks {
1759        let base = chunk * 4;
1760
1761        // Load 4 bytes as a single i32 and zero-extend u8→u32 via SSE4.1
1762        let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
1763            input.as_ptr().add(base) as *const i32
1764        ));
1765        let ints = _mm_cvtepu8_epi32(bytes);
1766        let floats = _mm_cvtepi32_ps(ints);
1767
1768        // Apply scale and min_val: result = floats * scale + min_val
1769        let scaled = _mm_add_ps(_mm_mul_ps(floats, scale_v), min_v);
1770
1771        _mm_storeu_ps(output.as_mut_ptr().add(base), scaled);
1772    }
1773
1774    // Handle remainder
1775    let base = chunks * 4;
1776    for i in 0..remainder {
1777        output[base + i] = input[base + i] as f32 * scale + min_val;
1778    }
1779}
1780
1781// ============================================================================
1782// Algebraic float reductions
1783// ============================================================================
1784//
1785// `f32::algebraic_add` / `algebraic_mul` (stable since Rust 1.98) permit the
1786// compiler to reassociate a floating-point reduction. That permission is the
1787// whole point: with strict IEEE `+` the loop-carried dependency on the
1788// accumulator pins these reductions to one scalar add per iteration, and LLVM
1789// is not allowed to split them into independent accumulator chains or vector
1790// lanes. With algebraic ops it vectorizes them the same way the hand-written
1791// NEON/AVX kernels below do by hand.
1792//
1793// Measured on aarch64 (Apple Silicon, rustc 1.98.0, opt-level=3, baseline
1794// target-cpu), before -> after ns/op:
1795//
1796//   dim    squared_l2       scalar dot     fused dot+norm   SOAR loss
1797//   128     55.5 -> 17.0     108 ->  16      145 ->  27      76.6 -> 18.8
1798//   384    220.8 -> 26.1     360 ->  32      365 ->  30      290  -> 80.0
1799//   768    592.8 -> 63.2     956 ->  55      752 ->  52      538  -> 67.7
1800//   1536  1827.8 -> 123.6   2032 -> 131     1404 -> 123     1464  -> 287
1801//
1802// i.e. 3-17x depending on width, largest at embedding-sized dimensions.
1803// Relative error against a strict f64 reference stays under 1e-6.
1804//
1805// End to end on `benches/vector_indexing.rs` (criterion, source-only diff):
1806// ivf_coarse_training/257 clusters -29.1%, /64 clusters -16.4%,
1807// ivf_tq_plan/64 -21.9%, ivf_tq_plan/16 -9.4%. See
1808// docs/algebraic-float-reductions.md.
1809//
1810// These operations are always safe (never UB), but they are *not*
1811// bit-reproducible across builds: a different rustc version, target CPU, or
1812// inlining decision may pick a different reduction order and move the last few
1813// ULPs. That variance already exists in every kernel in this module —
1814// `dot_product_f32` rounds differently on NEON (4 accumulators), AVX2 (4),
1815// AVX-512 (4) and the scalar path (1), so a query scored on an Apple Silicon
1816// replica already does not bit-match the same query on an AVX-512 replica.
1817// Using algebraic ops in the scalar paths therefore adds no new *class* of
1818// variance, only the same one at vector speed.
1819//
1820// Do not use these where a float is compared for bit-exact equality, hashed, or
1821// written into a content-addressed artifact.
1822
1823/// Dot product of two equal-length f32 slices, reassociated for vectorization.
1824#[inline]
1825fn dot_product_f32_scalar(a: &[f32], b: &[f32]) -> f32 {
1826    a.iter().zip(b).fold(0.0f32, |acc, (&x, &y)| {
1827        acc.algebraic_add(x.algebraic_mul(y))
1828    })
1829}
1830
1831/// Fused dot(a, b) and dot(b, b) over equal-length f32 slices in one pass.
1832#[inline]
1833fn fused_dot_norm_scalar(a: &[f32], b: &[f32]) -> (f32, f32) {
1834    a.iter()
1835        .zip(b)
1836        .fold((0.0f32, 0.0f32), |(dot, norm), (&x, &y)| {
1837            (
1838                dot.algebraic_add(x.algebraic_mul(y)),
1839                norm.algebraic_add(y.algebraic_mul(y)),
1840            )
1841        })
1842}
1843
1844/// Squared L2 distance `||a - b||^2` between two equal-length f32 slices.
1845///
1846/// The element-wise subtraction stays strict IEEE; only the summation is
1847/// reassociated. Iterates over `min(a.len(), b.len())` elements.
1848#[inline]
1849pub fn squared_l2_f32(a: &[f32], b: &[f32]) -> f32 {
1850    a.iter().zip(b).fold(0.0f32, |acc, (&x, &y)| {
1851        let delta = x - y;
1852        acc.algebraic_add(delta.algebraic_mul(delta))
1853    })
1854}
1855
1856/// Squared L2 norm `||v||^2` of an f32 slice.
1857#[inline]
1858pub fn norm_squared_f32(v: &[f32]) -> f32 {
1859    v.iter()
1860        .fold(0.0f32, |acc, &x| acc.algebraic_add(x.algebraic_mul(x)))
1861}
1862
1863/// L2 norm `||v||` of an f32 slice.
1864#[inline]
1865pub fn norm_f32(v: &[f32]) -> f32 {
1866    norm_squared_f32(v).sqrt()
1867}
1868
1869/// Compute dot product of two f32 arrays with SIMD acceleration
1870#[inline]
1871pub fn dot_product_f32(a: &[f32], b: &[f32], count: usize) -> f32 {
1872    assert!(
1873        count <= a.len() && count <= b.len(),
1874        "dot_product_f32 count {count} exceeds input lengths ({}, {})",
1875        a.len(),
1876        b.len()
1877    );
1878    #[cfg(target_arch = "aarch64")]
1879    {
1880        if neon::is_available() {
1881            return unsafe { dot_product_f32_neon(a, b, count) };
1882        }
1883    }
1884
1885    #[cfg(target_arch = "x86_64")]
1886    {
1887        if is_x86_feature_detected!("avx512f") {
1888            return unsafe { dot_product_f32_avx512(a, b, count) };
1889        }
1890        if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
1891            return unsafe { dot_product_f32_avx2(a, b, count) };
1892        }
1893        if sse::is_available() {
1894            return unsafe { dot_product_f32_sse(a, b, count) };
1895        }
1896    }
1897
1898    // Scalar fallback
1899    dot_product_f32_scalar(&a[..count], &b[..count])
1900}
1901
1902#[cfg(target_arch = "aarch64")]
1903#[target_feature(enable = "neon")]
1904#[allow(unsafe_op_in_unsafe_fn)]
1905unsafe fn dot_product_f32_neon(a: &[f32], b: &[f32], count: usize) -> f32 {
1906    use std::arch::aarch64::*;
1907
1908    let chunks16 = count / 16;
1909    let remainder = count % 16;
1910
1911    let mut acc0 = vdupq_n_f32(0.0);
1912    let mut acc1 = vdupq_n_f32(0.0);
1913    let mut acc2 = vdupq_n_f32(0.0);
1914    let mut acc3 = vdupq_n_f32(0.0);
1915
1916    for c in 0..chunks16 {
1917        let base = c * 16;
1918        acc0 = vfmaq_f32(
1919            acc0,
1920            vld1q_f32(a.as_ptr().add(base)),
1921            vld1q_f32(b.as_ptr().add(base)),
1922        );
1923        acc1 = vfmaq_f32(
1924            acc1,
1925            vld1q_f32(a.as_ptr().add(base + 4)),
1926            vld1q_f32(b.as_ptr().add(base + 4)),
1927        );
1928        acc2 = vfmaq_f32(
1929            acc2,
1930            vld1q_f32(a.as_ptr().add(base + 8)),
1931            vld1q_f32(b.as_ptr().add(base + 8)),
1932        );
1933        acc3 = vfmaq_f32(
1934            acc3,
1935            vld1q_f32(a.as_ptr().add(base + 12)),
1936            vld1q_f32(b.as_ptr().add(base + 12)),
1937        );
1938    }
1939
1940    let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
1941    let mut sum = vaddvq_f32(acc);
1942
1943    let base = chunks16 * 16;
1944    for i in 0..remainder {
1945        sum += a[base + i] * b[base + i];
1946    }
1947
1948    sum
1949}
1950
1951#[cfg(target_arch = "x86_64")]
1952#[target_feature(enable = "avx2", enable = "fma")]
1953#[allow(unsafe_op_in_unsafe_fn)]
1954unsafe fn dot_product_f32_avx2(a: &[f32], b: &[f32], count: usize) -> f32 {
1955    use std::arch::x86_64::*;
1956
1957    let chunks32 = count / 32;
1958    let remainder = count % 32;
1959
1960    let mut acc0 = _mm256_setzero_ps();
1961    let mut acc1 = _mm256_setzero_ps();
1962    let mut acc2 = _mm256_setzero_ps();
1963    let mut acc3 = _mm256_setzero_ps();
1964
1965    for c in 0..chunks32 {
1966        let base = c * 32;
1967        acc0 = _mm256_fmadd_ps(
1968            _mm256_loadu_ps(a.as_ptr().add(base)),
1969            _mm256_loadu_ps(b.as_ptr().add(base)),
1970            acc0,
1971        );
1972        acc1 = _mm256_fmadd_ps(
1973            _mm256_loadu_ps(a.as_ptr().add(base + 8)),
1974            _mm256_loadu_ps(b.as_ptr().add(base + 8)),
1975            acc1,
1976        );
1977        acc2 = _mm256_fmadd_ps(
1978            _mm256_loadu_ps(a.as_ptr().add(base + 16)),
1979            _mm256_loadu_ps(b.as_ptr().add(base + 16)),
1980            acc2,
1981        );
1982        acc3 = _mm256_fmadd_ps(
1983            _mm256_loadu_ps(a.as_ptr().add(base + 24)),
1984            _mm256_loadu_ps(b.as_ptr().add(base + 24)),
1985            acc3,
1986        );
1987    }
1988
1989    let acc = _mm256_add_ps(_mm256_add_ps(acc0, acc1), _mm256_add_ps(acc2, acc3));
1990
1991    // Horizontal sum: 256-bit → 128-bit → scalar
1992    let hi = _mm256_extractf128_ps(acc, 1);
1993    let lo = _mm256_castps256_ps128(acc);
1994    let sum128 = _mm_add_ps(lo, hi);
1995    let shuf = _mm_shuffle_ps(sum128, sum128, 0b10_11_00_01);
1996    let sums = _mm_add_ps(sum128, shuf);
1997    let shuf2 = _mm_movehl_ps(sums, sums);
1998    let final_sum = _mm_add_ss(sums, shuf2);
1999
2000    let mut sum = _mm_cvtss_f32(final_sum);
2001
2002    let base = chunks32 * 32;
2003    for i in 0..remainder {
2004        sum += a[base + i] * b[base + i];
2005    }
2006
2007    sum
2008}
2009
2010#[cfg(target_arch = "x86_64")]
2011#[target_feature(enable = "sse")]
2012#[allow(unsafe_op_in_unsafe_fn)]
2013unsafe fn dot_product_f32_sse(a: &[f32], b: &[f32], count: usize) -> f32 {
2014    use std::arch::x86_64::*;
2015
2016    let chunks = count / 4;
2017    let remainder = count % 4;
2018
2019    let mut acc = _mm_setzero_ps();
2020
2021    for chunk in 0..chunks {
2022        let base = chunk * 4;
2023        let va = _mm_loadu_ps(a.as_ptr().add(base));
2024        let vb = _mm_loadu_ps(b.as_ptr().add(base));
2025        acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
2026    }
2027
2028    // Horizontal sum: [a, b, c, d] -> a + b + c + d
2029    let shuf = _mm_shuffle_ps(acc, acc, 0b10_11_00_01); // [b, a, d, c]
2030    let sums = _mm_add_ps(acc, shuf); // [a+b, a+b, c+d, c+d]
2031    let shuf2 = _mm_movehl_ps(sums, sums); // [c+d, c+d, ?, ?]
2032    let final_sum = _mm_add_ss(sums, shuf2); // [a+b+c+d, ?, ?, ?]
2033
2034    let mut sum = _mm_cvtss_f32(final_sum);
2035
2036    // Handle remainder
2037    let base = chunks * 4;
2038    for i in 0..remainder {
2039        sum += a[base + i] * b[base + i];
2040    }
2041
2042    sum
2043}
2044
2045#[cfg(target_arch = "x86_64")]
2046#[target_feature(enable = "avx512f")]
2047#[allow(unsafe_op_in_unsafe_fn)]
2048unsafe fn dot_product_f32_avx512(a: &[f32], b: &[f32], count: usize) -> f32 {
2049    use std::arch::x86_64::*;
2050
2051    let chunks64 = count / 64;
2052    let remainder = count % 64;
2053
2054    let mut acc0 = _mm512_setzero_ps();
2055    let mut acc1 = _mm512_setzero_ps();
2056    let mut acc2 = _mm512_setzero_ps();
2057    let mut acc3 = _mm512_setzero_ps();
2058
2059    for c in 0..chunks64 {
2060        let base = c * 64;
2061        acc0 = _mm512_fmadd_ps(
2062            _mm512_loadu_ps(a.as_ptr().add(base)),
2063            _mm512_loadu_ps(b.as_ptr().add(base)),
2064            acc0,
2065        );
2066        acc1 = _mm512_fmadd_ps(
2067            _mm512_loadu_ps(a.as_ptr().add(base + 16)),
2068            _mm512_loadu_ps(b.as_ptr().add(base + 16)),
2069            acc1,
2070        );
2071        acc2 = _mm512_fmadd_ps(
2072            _mm512_loadu_ps(a.as_ptr().add(base + 32)),
2073            _mm512_loadu_ps(b.as_ptr().add(base + 32)),
2074            acc2,
2075        );
2076        acc3 = _mm512_fmadd_ps(
2077            _mm512_loadu_ps(a.as_ptr().add(base + 48)),
2078            _mm512_loadu_ps(b.as_ptr().add(base + 48)),
2079            acc3,
2080        );
2081    }
2082
2083    let acc = _mm512_add_ps(_mm512_add_ps(acc0, acc1), _mm512_add_ps(acc2, acc3));
2084    let mut sum = _mm512_reduce_add_ps(acc);
2085
2086    let base = chunks64 * 64;
2087    for i in 0..remainder {
2088        sum += a[base + i] * b[base + i];
2089    }
2090
2091    sum
2092}
2093
2094#[cfg(target_arch = "x86_64")]
2095#[target_feature(enable = "avx512f")]
2096#[allow(unsafe_op_in_unsafe_fn)]
2097unsafe fn fused_dot_norm_avx512(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2098    use std::arch::x86_64::*;
2099
2100    let chunks64 = count / 64;
2101    let remainder = count % 64;
2102
2103    let mut d0 = _mm512_setzero_ps();
2104    let mut d1 = _mm512_setzero_ps();
2105    let mut d2 = _mm512_setzero_ps();
2106    let mut d3 = _mm512_setzero_ps();
2107    let mut n0 = _mm512_setzero_ps();
2108    let mut n1 = _mm512_setzero_ps();
2109    let mut n2 = _mm512_setzero_ps();
2110    let mut n3 = _mm512_setzero_ps();
2111
2112    for c in 0..chunks64 {
2113        let base = c * 64;
2114        let vb0 = _mm512_loadu_ps(b.as_ptr().add(base));
2115        d0 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2116        n0 = _mm512_fmadd_ps(vb0, vb0, n0);
2117        let vb1 = _mm512_loadu_ps(b.as_ptr().add(base + 16));
2118        d1 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 16)), vb1, d1);
2119        n1 = _mm512_fmadd_ps(vb1, vb1, n1);
2120        let vb2 = _mm512_loadu_ps(b.as_ptr().add(base + 32));
2121        d2 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 32)), vb2, d2);
2122        n2 = _mm512_fmadd_ps(vb2, vb2, n2);
2123        let vb3 = _mm512_loadu_ps(b.as_ptr().add(base + 48));
2124        d3 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 48)), vb3, d3);
2125        n3 = _mm512_fmadd_ps(vb3, vb3, n3);
2126    }
2127
2128    let acc_dot = _mm512_add_ps(_mm512_add_ps(d0, d1), _mm512_add_ps(d2, d3));
2129    let acc_norm = _mm512_add_ps(_mm512_add_ps(n0, n1), _mm512_add_ps(n2, n3));
2130    let mut dot = _mm512_reduce_add_ps(acc_dot);
2131    let mut norm = _mm512_reduce_add_ps(acc_norm);
2132
2133    let base = chunks64 * 64;
2134    for i in 0..remainder {
2135        dot += a[base + i] * b[base + i];
2136        norm += b[base + i] * b[base + i];
2137    }
2138
2139    (dot, norm)
2140}
2141
2142// ============================================================================
2143// Batched Cosine Similarity for Dense Vector Search
2144// ============================================================================
2145
2146/// Fused dot-product + self-norm in a single pass (SIMD accelerated).
2147///
2148/// Returns (dot(a, b), dot(b, b)) — i.e. the dot product of a·b and ||b||².
2149/// Loads `b` only once (halves memory bandwidth vs two separate dot products).
2150#[inline]
2151fn fused_dot_norm(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2152    #[cfg(target_arch = "aarch64")]
2153    {
2154        if neon::is_available() {
2155            return unsafe { fused_dot_norm_neon(a, b, count) };
2156        }
2157    }
2158
2159    #[cfg(target_arch = "x86_64")]
2160    {
2161        if is_x86_feature_detected!("avx512f") {
2162            return unsafe { fused_dot_norm_avx512(a, b, count) };
2163        }
2164        if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
2165            return unsafe { fused_dot_norm_avx2(a, b, count) };
2166        }
2167        if sse::is_available() {
2168            return unsafe { fused_dot_norm_sse(a, b, count) };
2169        }
2170    }
2171
2172    // Scalar fallback
2173    fused_dot_norm_scalar(&a[..count], &b[..count])
2174}
2175
2176#[cfg(target_arch = "aarch64")]
2177#[target_feature(enable = "neon")]
2178#[allow(unsafe_op_in_unsafe_fn)]
2179unsafe fn fused_dot_norm_neon(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2180    use std::arch::aarch64::*;
2181
2182    let chunks16 = count / 16;
2183    let remainder = count % 16;
2184
2185    let mut d0 = vdupq_n_f32(0.0);
2186    let mut d1 = vdupq_n_f32(0.0);
2187    let mut d2 = vdupq_n_f32(0.0);
2188    let mut d3 = vdupq_n_f32(0.0);
2189    let mut n0 = vdupq_n_f32(0.0);
2190    let mut n1 = vdupq_n_f32(0.0);
2191    let mut n2 = vdupq_n_f32(0.0);
2192    let mut n3 = vdupq_n_f32(0.0);
2193
2194    for c in 0..chunks16 {
2195        let base = c * 16;
2196        let va0 = vld1q_f32(a.as_ptr().add(base));
2197        let vb0 = vld1q_f32(b.as_ptr().add(base));
2198        d0 = vfmaq_f32(d0, va0, vb0);
2199        n0 = vfmaq_f32(n0, vb0, vb0);
2200        let va1 = vld1q_f32(a.as_ptr().add(base + 4));
2201        let vb1 = vld1q_f32(b.as_ptr().add(base + 4));
2202        d1 = vfmaq_f32(d1, va1, vb1);
2203        n1 = vfmaq_f32(n1, vb1, vb1);
2204        let va2 = vld1q_f32(a.as_ptr().add(base + 8));
2205        let vb2 = vld1q_f32(b.as_ptr().add(base + 8));
2206        d2 = vfmaq_f32(d2, va2, vb2);
2207        n2 = vfmaq_f32(n2, vb2, vb2);
2208        let va3 = vld1q_f32(a.as_ptr().add(base + 12));
2209        let vb3 = vld1q_f32(b.as_ptr().add(base + 12));
2210        d3 = vfmaq_f32(d3, va3, vb3);
2211        n3 = vfmaq_f32(n3, vb3, vb3);
2212    }
2213
2214    let acc_dot = vaddq_f32(vaddq_f32(d0, d1), vaddq_f32(d2, d3));
2215    let acc_norm = vaddq_f32(vaddq_f32(n0, n1), vaddq_f32(n2, n3));
2216    let mut dot = vaddvq_f32(acc_dot);
2217    let mut norm = vaddvq_f32(acc_norm);
2218
2219    let base = chunks16 * 16;
2220    for i in 0..remainder {
2221        dot += a[base + i] * b[base + i];
2222        norm += b[base + i] * b[base + i];
2223    }
2224
2225    (dot, norm)
2226}
2227
2228#[cfg(target_arch = "x86_64")]
2229#[target_feature(enable = "avx2", enable = "fma")]
2230#[allow(unsafe_op_in_unsafe_fn)]
2231unsafe fn fused_dot_norm_avx2(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2232    use std::arch::x86_64::*;
2233
2234    let chunks32 = count / 32;
2235    let remainder = count % 32;
2236
2237    let mut d0 = _mm256_setzero_ps();
2238    let mut d1 = _mm256_setzero_ps();
2239    let mut d2 = _mm256_setzero_ps();
2240    let mut d3 = _mm256_setzero_ps();
2241    let mut n0 = _mm256_setzero_ps();
2242    let mut n1 = _mm256_setzero_ps();
2243    let mut n2 = _mm256_setzero_ps();
2244    let mut n3 = _mm256_setzero_ps();
2245
2246    for c in 0..chunks32 {
2247        let base = c * 32;
2248        let vb0 = _mm256_loadu_ps(b.as_ptr().add(base));
2249        d0 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2250        n0 = _mm256_fmadd_ps(vb0, vb0, n0);
2251        let vb1 = _mm256_loadu_ps(b.as_ptr().add(base + 8));
2252        d1 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 8)), vb1, d1);
2253        n1 = _mm256_fmadd_ps(vb1, vb1, n1);
2254        let vb2 = _mm256_loadu_ps(b.as_ptr().add(base + 16));
2255        d2 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 16)), vb2, d2);
2256        n2 = _mm256_fmadd_ps(vb2, vb2, n2);
2257        let vb3 = _mm256_loadu_ps(b.as_ptr().add(base + 24));
2258        d3 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 24)), vb3, d3);
2259        n3 = _mm256_fmadd_ps(vb3, vb3, n3);
2260    }
2261
2262    let acc_dot = _mm256_add_ps(_mm256_add_ps(d0, d1), _mm256_add_ps(d2, d3));
2263    let acc_norm = _mm256_add_ps(_mm256_add_ps(n0, n1), _mm256_add_ps(n2, n3));
2264
2265    // Horizontal sums: 256→128→scalar
2266    let hi_d = _mm256_extractf128_ps(acc_dot, 1);
2267    let lo_d = _mm256_castps256_ps128(acc_dot);
2268    let sum_d = _mm_add_ps(lo_d, hi_d);
2269    let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
2270    let sums_d = _mm_add_ps(sum_d, shuf_d);
2271    let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2272    let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2273
2274    let hi_n = _mm256_extractf128_ps(acc_norm, 1);
2275    let lo_n = _mm256_castps256_ps128(acc_norm);
2276    let sum_n = _mm_add_ps(lo_n, hi_n);
2277    let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
2278    let sums_n = _mm_add_ps(sum_n, shuf_n);
2279    let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2280    let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2281
2282    let base = chunks32 * 32;
2283    for i in 0..remainder {
2284        dot += a[base + i] * b[base + i];
2285        norm += b[base + i] * b[base + i];
2286    }
2287
2288    (dot, norm)
2289}
2290
2291#[cfg(target_arch = "x86_64")]
2292#[target_feature(enable = "sse")]
2293#[allow(unsafe_op_in_unsafe_fn)]
2294unsafe fn fused_dot_norm_sse(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2295    use std::arch::x86_64::*;
2296
2297    let chunks = count / 4;
2298    let remainder = count % 4;
2299
2300    let mut acc_dot = _mm_setzero_ps();
2301    let mut acc_norm = _mm_setzero_ps();
2302
2303    for chunk in 0..chunks {
2304        let base = chunk * 4;
2305        let va = _mm_loadu_ps(a.as_ptr().add(base));
2306        let vb = _mm_loadu_ps(b.as_ptr().add(base));
2307        acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2308        acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2309    }
2310
2311    // Horizontal sums
2312    let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2313    let sums_d = _mm_add_ps(acc_dot, shuf_d);
2314    let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2315    let final_d = _mm_add_ss(sums_d, shuf2_d);
2316    let mut dot = _mm_cvtss_f32(final_d);
2317
2318    let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2319    let sums_n = _mm_add_ps(acc_norm, shuf_n);
2320    let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2321    let final_n = _mm_add_ss(sums_n, shuf2_n);
2322    let mut norm = _mm_cvtss_f32(final_n);
2323
2324    let base = chunks * 4;
2325    for i in 0..remainder {
2326        dot += a[base + i] * b[base + i];
2327        norm += b[base + i] * b[base + i];
2328    }
2329
2330    (dot, norm)
2331}
2332
2333/// Fast approximate reciprocal square root: 1/sqrt(x).
2334///
2335/// Uses the IEEE 754 bit trick (Quake III) + one Newton-Raphson iteration
2336/// for ~23-bit precision — sufficient for cosine similarity scoring.
2337/// ~3-5x faster than `1.0 / x.sqrt()` on most architectures.
2338#[inline]
2339pub fn fast_inv_sqrt(x: f32) -> f32 {
2340    let half = 0.5 * x;
2341    let i = 0x5F37_5A86_u32.wrapping_sub(x.to_bits() >> 1);
2342    let y = f32::from_bits(i);
2343    let y = y * (1.5 - half * y * y); // first Newton-Raphson step
2344    y * (1.5 - half * y * y) // second step: ~23-bit precision
2345}
2346
2347/// Batch cosine similarity: query vs N contiguous vectors.
2348///
2349/// `vectors` is a contiguous buffer of `n * dim` floats (row-major).
2350/// `scores` must have length >= n.
2351///
2352/// Optimizations over calling `cosine_similarity` N times:
2353/// 1. Query norm computed once (not N times)
2354/// 2. Fused dot+norm kernel — each vector loaded once (halves bandwidth)
2355/// 3. No per-call overhead (branch prediction, function calls)
2356/// 4. Fast reciprocal square root (~3-5x faster than 1/sqrt)
2357#[inline]
2358pub fn batch_cosine_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
2359    let n = scores.len();
2360    let required = n
2361        .checked_mul(dim)
2362        .expect("batch cosine vector length overflow");
2363    assert_eq!(query.len(), dim, "batch cosine query dimension mismatch");
2364    assert!(
2365        vectors.len() >= required,
2366        "batch cosine vectors are truncated: need {required}, got {}",
2367        vectors.len()
2368    );
2369
2370    if dim == 0 || n == 0 {
2371        return;
2372    }
2373
2374    // Pre-compute query inverse norm once
2375    let norm_q_sq = dot_product_f32(query, query, dim);
2376    if norm_q_sq < f32::EPSILON {
2377        for s in scores.iter_mut() {
2378            *s = 0.0;
2379        }
2380        return;
2381    }
2382    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
2383
2384    for i in 0..n {
2385        let vec = &vectors[i * dim..(i + 1) * dim];
2386        let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
2387        if norm_v_sq < f32::EPSILON {
2388            scores[i] = 0.0;
2389        } else {
2390            scores[i] = dot * inv_norm_q * fast_inv_sqrt(norm_v_sq);
2391        }
2392    }
2393}
2394
2395// ============================================================================
2396// f16 (IEEE 754 half-precision) conversion
2397// ============================================================================
2398
2399/// Convert f32 to f16 (IEEE 754 half-precision), stored as u16
2400#[inline]
2401pub fn f32_to_f16(value: f32) -> u16 {
2402    let bits = value.to_bits();
2403    let sign = (bits >> 16) & 0x8000;
2404    let exp = ((bits >> 23) & 0xFF) as i32;
2405    let mantissa = bits & 0x7F_FFFF;
2406
2407    if exp == 255 {
2408        // Inf/NaN
2409        return (sign | 0x7C00 | ((mantissa >> 13) & 0x3FF)) as u16;
2410    }
2411
2412    let exp16 = exp - 127 + 15;
2413
2414    if exp16 >= 31 {
2415        return (sign | 0x7C00) as u16; // overflow → infinity
2416    }
2417
2418    if exp16 <= 0 {
2419        if exp16 < -10 {
2420            return sign as u16; // too small → zero
2421        }
2422        let shift = (1 - exp16) as u32;
2423        let m = (mantissa | 0x80_0000) >> shift;
2424        // Round-to-nearest-even
2425        let round_bit = (m >> 12) & 1;
2426        let sticky = m & 0xFFF;
2427        let m13 = m >> 13;
2428        let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2429        return (sign | rounded) as u16;
2430    }
2431
2432    // Round-to-nearest-even for normal numbers
2433    let round_bit = (mantissa >> 12) & 1;
2434    let sticky = mantissa & 0xFFF;
2435    let m13 = mantissa >> 13;
2436    let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2437    // Check if rounding caused mantissa overflow (carry into exponent)
2438    if rounded > 0x3FF {
2439        let exp16_inc = exp16 as u32 + 1;
2440        if exp16_inc >= 31 {
2441            return (sign | 0x7C00) as u16; // overflow → infinity
2442        }
2443        (sign | (exp16_inc << 10)) as u16
2444    } else {
2445        (sign | ((exp16 as u32) << 10) | rounded) as u16
2446    }
2447}
2448
2449/// Convert f16 (stored as u16) to f32
2450#[inline]
2451pub fn f16_to_f32(half: u16) -> f32 {
2452    let sign = ((half & 0x8000) as u32) << 16;
2453    let exp = ((half >> 10) & 0x1F) as u32;
2454    let mantissa = (half & 0x3FF) as u32;
2455
2456    if exp == 0 {
2457        if mantissa == 0 {
2458            return f32::from_bits(sign);
2459        }
2460        // Subnormal: normalize
2461        let mut e = 0u32;
2462        let mut m = mantissa;
2463        while (m & 0x400) == 0 {
2464            m <<= 1;
2465            e += 1;
2466        }
2467        return f32::from_bits(sign | ((127 - 15 + 1 - e) << 23) | ((m & 0x3FF) << 13));
2468    }
2469
2470    if exp == 31 {
2471        return f32::from_bits(sign | 0x7F80_0000 | (mantissa << 13));
2472    }
2473
2474    f32::from_bits(sign | ((exp + 127 - 15) << 23) | (mantissa << 13))
2475}
2476
2477// ============================================================================
2478// uint8 scalar quantization for [-1, 1] range
2479// ============================================================================
2480
2481const U8_SCALE: f32 = 127.5;
2482const U8_INV_SCALE: f32 = 1.0 / 127.5;
2483
2484/// Quantize f32 in [-1, 1] to u8 [0, 255]
2485#[inline]
2486pub fn f32_to_u8_saturating(value: f32) -> u8 {
2487    ((value.clamp(-1.0, 1.0) + 1.0) * U8_SCALE) as u8
2488}
2489
2490/// Dequantize u8 [0, 255] to f32 in [-1, 1]
2491#[inline]
2492pub fn u8_to_f32(byte: u8) -> f32 {
2493    byte as f32 * U8_INV_SCALE - 1.0
2494}
2495
2496// ============================================================================
2497// Batch conversion (used during builder write)
2498// ============================================================================
2499
2500/// Batch convert f32 slice to f16 (stored as u16)
2501pub fn batch_f32_to_f16(src: &[f32], dst: &mut [u16]) {
2502    debug_assert_eq!(src.len(), dst.len());
2503    for (s, d) in src.iter().zip(dst.iter_mut()) {
2504        *d = f32_to_f16(*s);
2505    }
2506}
2507
2508/// Batch convert an f32 slice to u8 with `[-1, 1]` to `[0, 255]` mapping.
2509pub fn batch_f32_to_u8(src: &[f32], dst: &mut [u8]) {
2510    debug_assert_eq!(src.len(), dst.len());
2511    for (s, d) in src.iter().zip(dst.iter_mut()) {
2512        *d = f32_to_u8_saturating(*s);
2513    }
2514}
2515
2516// ============================================================================
2517// NEON-accelerated fused dot+norm for quantized vectors
2518// ============================================================================
2519
2520#[cfg(target_arch = "aarch64")]
2521#[allow(unsafe_op_in_unsafe_fn)]
2522mod neon_quant {
2523    use std::arch::aarch64::*;
2524
2525    /// Fused dot(query_f16, vec_f16) + norm(vec_f16) for f16 vectors on NEON.
2526    ///
2527    /// Both query and vectors are f16 (stored as u16). Uses hardware `vcvt_f32_f16`
2528    /// for SIMD f16→f32 conversion (replaces scalar bit manipulation), processes
2529    /// 8 elements per iteration with f32 accumulation for precision.
2530    #[allow(clippy::incompatible_msrv)]
2531    #[target_feature(enable = "neon")]
2532    pub unsafe fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2533        let chunks16 = dim / 16;
2534        let remainder = dim % 16;
2535
2536        // 2 accumulator pairs to hide FMA latency (processes 16 f16 per iteration)
2537        let mut acc_dot0 = vdupq_n_f32(0.0);
2538        let mut acc_dot1 = vdupq_n_f32(0.0);
2539        let mut acc_norm0 = vdupq_n_f32(0.0);
2540        let mut acc_norm1 = vdupq_n_f32(0.0);
2541
2542        for c in 0..chunks16 {
2543            let base = c * 16;
2544
2545            // First 8 f16 elements
2546            let v_raw0 = vld1q_u16(vec_f16.as_ptr().add(base));
2547            let v_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw0)));
2548            let v_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw0)));
2549            let q_raw0 = vld1q_u16(query_f16.as_ptr().add(base));
2550            let q_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw0)));
2551            let q_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw0)));
2552
2553            acc_dot0 = vfmaq_f32(acc_dot0, q_lo0, v_lo0);
2554            acc_dot0 = vfmaq_f32(acc_dot0, q_hi0, v_hi0);
2555            acc_norm0 = vfmaq_f32(acc_norm0, v_lo0, v_lo0);
2556            acc_norm0 = vfmaq_f32(acc_norm0, v_hi0, v_hi0);
2557
2558            // Second 8 f16 elements (independent accumulator chain)
2559            let v_raw1 = vld1q_u16(vec_f16.as_ptr().add(base + 8));
2560            let v_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw1)));
2561            let v_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw1)));
2562            let q_raw1 = vld1q_u16(query_f16.as_ptr().add(base + 8));
2563            let q_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw1)));
2564            let q_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw1)));
2565
2566            acc_dot1 = vfmaq_f32(acc_dot1, q_lo1, v_lo1);
2567            acc_dot1 = vfmaq_f32(acc_dot1, q_hi1, v_hi1);
2568            acc_norm1 = vfmaq_f32(acc_norm1, v_lo1, v_lo1);
2569            acc_norm1 = vfmaq_f32(acc_norm1, v_hi1, v_hi1);
2570        }
2571
2572        // Combine accumulator pairs
2573        let mut dot = vaddvq_f32(vaddq_f32(acc_dot0, acc_dot1));
2574        let mut norm = vaddvq_f32(vaddq_f32(acc_norm0, acc_norm1));
2575
2576        // Handle remainder
2577        let base = chunks16 * 16;
2578        for i in 0..remainder {
2579            let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
2580            let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
2581            dot += q * v;
2582            norm += v * v;
2583        }
2584
2585        (dot, norm)
2586    }
2587
2588    /// Fused dot(query, vec) + norm(vec) for u8 vectors on NEON.
2589    /// Processes 16 u8 values per iteration using NEON widening chain.
2590    #[target_feature(enable = "neon")]
2591    pub unsafe fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2592        let scale = vdupq_n_f32(super::U8_INV_SCALE);
2593        let offset = vdupq_n_f32(-1.0);
2594
2595        let chunks16 = dim / 16;
2596        let remainder = dim % 16;
2597
2598        let mut acc_dot = vdupq_n_f32(0.0);
2599        let mut acc_norm = vdupq_n_f32(0.0);
2600
2601        for c in 0..chunks16 {
2602            let base = c * 16;
2603
2604            // Load 16 u8 values
2605            let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
2606
2607            // Widen: 16×u8 → 2×8×u16 → 4×4×u32 → 4×4×f32
2608            let lo8 = vget_low_u8(bytes);
2609            let hi8 = vget_high_u8(bytes);
2610            let lo16 = vmovl_u8(lo8);
2611            let hi16 = vmovl_u8(hi8);
2612
2613            let f0 = vaddq_f32(
2614                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
2615                offset,
2616            );
2617            let f1 = vaddq_f32(
2618                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
2619                offset,
2620            );
2621            let f2 = vaddq_f32(
2622                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
2623                offset,
2624            );
2625            let f3 = vaddq_f32(
2626                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
2627                offset,
2628            );
2629
2630            let q0 = vld1q_f32(query.as_ptr().add(base));
2631            let q1 = vld1q_f32(query.as_ptr().add(base + 4));
2632            let q2 = vld1q_f32(query.as_ptr().add(base + 8));
2633            let q3 = vld1q_f32(query.as_ptr().add(base + 12));
2634
2635            acc_dot = vfmaq_f32(acc_dot, q0, f0);
2636            acc_dot = vfmaq_f32(acc_dot, q1, f1);
2637            acc_dot = vfmaq_f32(acc_dot, q2, f2);
2638            acc_dot = vfmaq_f32(acc_dot, q3, f3);
2639
2640            acc_norm = vfmaq_f32(acc_norm, f0, f0);
2641            acc_norm = vfmaq_f32(acc_norm, f1, f1);
2642            acc_norm = vfmaq_f32(acc_norm, f2, f2);
2643            acc_norm = vfmaq_f32(acc_norm, f3, f3);
2644        }
2645
2646        let mut dot = vaddvq_f32(acc_dot);
2647        let mut norm = vaddvq_f32(acc_norm);
2648
2649        let base = chunks16 * 16;
2650        for i in 0..remainder {
2651            let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
2652            dot += *query.get_unchecked(base + i) * v;
2653            norm += v * v;
2654        }
2655
2656        (dot, norm)
2657    }
2658
2659    /// Dot product only for f16 vectors on NEON (no norm — for unit_norm vectors).
2660    #[allow(clippy::incompatible_msrv)]
2661    #[target_feature(enable = "neon")]
2662    pub unsafe fn dot_product_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2663        let chunks8 = dim / 8;
2664        let remainder = dim % 8;
2665
2666        let mut acc = vdupq_n_f32(0.0);
2667
2668        for c in 0..chunks8 {
2669            let base = c * 8;
2670            let v_raw = vld1q_u16(vec_f16.as_ptr().add(base));
2671            let v_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw)));
2672            let v_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw)));
2673            let q_raw = vld1q_u16(query_f16.as_ptr().add(base));
2674            let q_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw)));
2675            let q_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw)));
2676            acc = vfmaq_f32(acc, q_lo, v_lo);
2677            acc = vfmaq_f32(acc, q_hi, v_hi);
2678        }
2679
2680        let mut dot = vaddvq_f32(acc);
2681        let base = chunks8 * 8;
2682        for i in 0..remainder {
2683            let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
2684            let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
2685            dot += q * v;
2686        }
2687        dot
2688    }
2689
2690    /// Dot product only for u8 vectors on NEON (no norm — for unit_norm vectors).
2691    #[target_feature(enable = "neon")]
2692    pub unsafe fn dot_product_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
2693        let scale = vdupq_n_f32(super::U8_INV_SCALE);
2694        let offset = vdupq_n_f32(-1.0);
2695        let chunks16 = dim / 16;
2696        let remainder = dim % 16;
2697
2698        let mut acc = vdupq_n_f32(0.0);
2699
2700        for c in 0..chunks16 {
2701            let base = c * 16;
2702            let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
2703            let lo8 = vget_low_u8(bytes);
2704            let hi8 = vget_high_u8(bytes);
2705            let lo16 = vmovl_u8(lo8);
2706            let hi16 = vmovl_u8(hi8);
2707            let f0 = vaddq_f32(
2708                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
2709                offset,
2710            );
2711            let f1 = vaddq_f32(
2712                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
2713                offset,
2714            );
2715            let f2 = vaddq_f32(
2716                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
2717                offset,
2718            );
2719            let f3 = vaddq_f32(
2720                vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
2721                offset,
2722            );
2723            let q0 = vld1q_f32(query.as_ptr().add(base));
2724            let q1 = vld1q_f32(query.as_ptr().add(base + 4));
2725            let q2 = vld1q_f32(query.as_ptr().add(base + 8));
2726            let q3 = vld1q_f32(query.as_ptr().add(base + 12));
2727            acc = vfmaq_f32(acc, q0, f0);
2728            acc = vfmaq_f32(acc, q1, f1);
2729            acc = vfmaq_f32(acc, q2, f2);
2730            acc = vfmaq_f32(acc, q3, f3);
2731        }
2732
2733        let mut dot = vaddvq_f32(acc);
2734        let base = chunks16 * 16;
2735        for i in 0..remainder {
2736            let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
2737            dot += *query.get_unchecked(base + i) * v;
2738        }
2739        dot
2740    }
2741}
2742
2743// ============================================================================
2744// Scalar fallback for fused dot+norm on quantized vectors
2745// ============================================================================
2746
2747#[allow(dead_code)]
2748fn fused_dot_norm_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2749    (0..dim).fold((0.0f32, 0.0f32), |(dot, norm), i| {
2750        let v = f16_to_f32(vec_f16[i]);
2751        let q = f16_to_f32(query_f16[i]);
2752        (
2753            dot.algebraic_add(q.algebraic_mul(v)),
2754            norm.algebraic_add(v.algebraic_mul(v)),
2755        )
2756    })
2757}
2758
2759#[allow(dead_code)]
2760fn fused_dot_norm_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2761    (0..dim).fold((0.0f32, 0.0f32), |(dot, norm), i| {
2762        let v = u8_to_f32(vec_u8[i]);
2763        (
2764            dot.algebraic_add(query[i].algebraic_mul(v)),
2765            norm.algebraic_add(v.algebraic_mul(v)),
2766        )
2767    })
2768}
2769
2770#[allow(dead_code)]
2771fn dot_product_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2772    (0..dim).fold(0.0f32, |dot, i| {
2773        dot.algebraic_add(f16_to_f32(query_f16[i]).algebraic_mul(f16_to_f32(vec_f16[i])))
2774    })
2775}
2776
2777#[allow(dead_code)]
2778fn dot_product_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
2779    (0..dim).fold(0.0f32, |dot, i| {
2780        dot.algebraic_add(query[i].algebraic_mul(u8_to_f32(vec_u8[i])))
2781    })
2782}
2783
2784// ============================================================================
2785// x86_64 SSE4.1 quantized fused dot+norm
2786// ============================================================================
2787
2788#[cfg(target_arch = "x86_64")]
2789#[target_feature(enable = "sse2", enable = "sse4.1")]
2790#[allow(unsafe_op_in_unsafe_fn)]
2791unsafe fn fused_dot_norm_f16_sse(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2792    use std::arch::x86_64::*;
2793
2794    let chunks = dim / 4;
2795    let remainder = dim % 4;
2796
2797    let mut acc_dot = _mm_setzero_ps();
2798    let mut acc_norm = _mm_setzero_ps();
2799
2800    for chunk in 0..chunks {
2801        let base = chunk * 4;
2802        // Load 4 f16 values and convert to f32 using scalar conversion
2803        let v0 = f16_to_f32(*vec_f16.get_unchecked(base));
2804        let v1 = f16_to_f32(*vec_f16.get_unchecked(base + 1));
2805        let v2 = f16_to_f32(*vec_f16.get_unchecked(base + 2));
2806        let v3 = f16_to_f32(*vec_f16.get_unchecked(base + 3));
2807        let vb = _mm_set_ps(v3, v2, v1, v0);
2808
2809        let q0 = f16_to_f32(*query_f16.get_unchecked(base));
2810        let q1 = f16_to_f32(*query_f16.get_unchecked(base + 1));
2811        let q2 = f16_to_f32(*query_f16.get_unchecked(base + 2));
2812        let q3 = f16_to_f32(*query_f16.get_unchecked(base + 3));
2813        let va = _mm_set_ps(q3, q2, q1, q0);
2814
2815        acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2816        acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2817    }
2818
2819    // Horizontal sums
2820    let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2821    let sums_d = _mm_add_ps(acc_dot, shuf_d);
2822    let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2823    let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2824
2825    let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2826    let sums_n = _mm_add_ps(acc_norm, shuf_n);
2827    let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2828    let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2829
2830    let base = chunks * 4;
2831    for i in 0..remainder {
2832        let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2833        let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2834        dot += q * v;
2835        norm += v * v;
2836    }
2837
2838    (dot, norm)
2839}
2840
2841#[cfg(target_arch = "x86_64")]
2842#[target_feature(enable = "sse2", enable = "sse4.1")]
2843#[allow(unsafe_op_in_unsafe_fn)]
2844unsafe fn fused_dot_norm_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2845    use std::arch::x86_64::*;
2846
2847    let scale = _mm_set1_ps(U8_INV_SCALE);
2848    let offset = _mm_set1_ps(-1.0);
2849
2850    let chunks = dim / 4;
2851    let remainder = dim % 4;
2852
2853    let mut acc_dot = _mm_setzero_ps();
2854    let mut acc_norm = _mm_setzero_ps();
2855
2856    for chunk in 0..chunks {
2857        let base = chunk * 4;
2858
2859        // Load 4 bytes, zero-extend to i32, convert to f32, dequantize
2860        let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
2861            vec_u8.as_ptr().add(base) as *const i32
2862        ));
2863        let ints = _mm_cvtepu8_epi32(bytes);
2864        let floats = _mm_cvtepi32_ps(ints);
2865        let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
2866
2867        let va = _mm_loadu_ps(query.as_ptr().add(base));
2868
2869        acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2870        acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2871    }
2872
2873    // Horizontal sums
2874    let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2875    let sums_d = _mm_add_ps(acc_dot, shuf_d);
2876    let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2877    let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2878
2879    let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2880    let sums_n = _mm_add_ps(acc_norm, shuf_n);
2881    let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2882    let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2883
2884    let base = chunks * 4;
2885    for i in 0..remainder {
2886        let v = u8_to_f32(*vec_u8.get_unchecked(base + i));
2887        dot += *query.get_unchecked(base + i) * v;
2888        norm += v * v;
2889    }
2890
2891    (dot, norm)
2892}
2893
2894// ============================================================================
2895// x86_64 F16C + AVX + FMA accelerated f16 scoring
2896// ============================================================================
2897
2898#[cfg(target_arch = "x86_64")]
2899#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
2900#[allow(unsafe_op_in_unsafe_fn)]
2901unsafe fn fused_dot_norm_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2902    use std::arch::x86_64::*;
2903
2904    let chunks16 = dim / 16;
2905    let remainder = dim % 16;
2906
2907    // 2 accumulator pairs to hide FMA latency (processes 16 f16 per iteration)
2908    let mut acc_dot0 = _mm256_setzero_ps();
2909    let mut acc_dot1 = _mm256_setzero_ps();
2910    let mut acc_norm0 = _mm256_setzero_ps();
2911    let mut acc_norm1 = _mm256_setzero_ps();
2912
2913    for c in 0..chunks16 {
2914        let base = c * 16;
2915
2916        // First 8 f16 elements
2917        let v_raw0 = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
2918        let vb0 = _mm256_cvtph_ps(v_raw0);
2919        let q_raw0 = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
2920        let qa0 = _mm256_cvtph_ps(q_raw0);
2921        acc_dot0 = _mm256_fmadd_ps(qa0, vb0, acc_dot0);
2922        acc_norm0 = _mm256_fmadd_ps(vb0, vb0, acc_norm0);
2923
2924        // Second 8 f16 elements (independent accumulator chain)
2925        let v_raw1 = _mm_loadu_si128(vec_f16.as_ptr().add(base + 8) as *const __m128i);
2926        let vb1 = _mm256_cvtph_ps(v_raw1);
2927        let q_raw1 = _mm_loadu_si128(query_f16.as_ptr().add(base + 8) as *const __m128i);
2928        let qa1 = _mm256_cvtph_ps(q_raw1);
2929        acc_dot1 = _mm256_fmadd_ps(qa1, vb1, acc_dot1);
2930        acc_norm1 = _mm256_fmadd_ps(vb1, vb1, acc_norm1);
2931    }
2932
2933    // Combine accumulator pairs
2934    let acc_dot = _mm256_add_ps(acc_dot0, acc_dot1);
2935    let acc_norm = _mm256_add_ps(acc_norm0, acc_norm1);
2936
2937    // Horizontal sum 256→128→scalar
2938    let hi_d = _mm256_extractf128_ps(acc_dot, 1);
2939    let lo_d = _mm256_castps256_ps128(acc_dot);
2940    let sum_d = _mm_add_ps(lo_d, hi_d);
2941    let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
2942    let sums_d = _mm_add_ps(sum_d, shuf_d);
2943    let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2944    let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2945
2946    let hi_n = _mm256_extractf128_ps(acc_norm, 1);
2947    let lo_n = _mm256_castps256_ps128(acc_norm);
2948    let sum_n = _mm_add_ps(lo_n, hi_n);
2949    let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
2950    let sums_n = _mm_add_ps(sum_n, shuf_n);
2951    let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2952    let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2953
2954    let base = chunks16 * 16;
2955    for i in 0..remainder {
2956        let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2957        let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2958        dot += q * v;
2959        norm += v * v;
2960    }
2961
2962    (dot, norm)
2963}
2964
2965#[cfg(target_arch = "x86_64")]
2966#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
2967#[allow(unsafe_op_in_unsafe_fn)]
2968unsafe fn dot_product_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2969    use std::arch::x86_64::*;
2970
2971    let chunks = dim / 8;
2972    let remainder = dim % 8;
2973    let mut acc = _mm256_setzero_ps();
2974
2975    for chunk in 0..chunks {
2976        let base = chunk * 8;
2977        let v_raw = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
2978        let vb = _mm256_cvtph_ps(v_raw);
2979        let q_raw = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
2980        let qa = _mm256_cvtph_ps(q_raw);
2981        acc = _mm256_fmadd_ps(qa, vb, acc);
2982    }
2983
2984    let hi = _mm256_extractf128_ps(acc, 1);
2985    let lo = _mm256_castps256_ps128(acc);
2986    let sum = _mm_add_ps(lo, hi);
2987    let shuf = _mm_shuffle_ps(sum, sum, 0b10_11_00_01);
2988    let sums = _mm_add_ps(sum, shuf);
2989    let shuf2 = _mm_movehl_ps(sums, sums);
2990    let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
2991
2992    let base = chunks * 8;
2993    for i in 0..remainder {
2994        let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2995        let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2996        dot += q * v;
2997    }
2998    dot
2999}
3000
3001#[cfg(target_arch = "x86_64")]
3002#[target_feature(enable = "sse2", enable = "sse4.1")]
3003#[allow(unsafe_op_in_unsafe_fn)]
3004unsafe fn dot_product_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3005    use std::arch::x86_64::*;
3006
3007    let scale = _mm_set1_ps(U8_INV_SCALE);
3008    let offset = _mm_set1_ps(-1.0);
3009    let chunks = dim / 4;
3010    let remainder = dim % 4;
3011    let mut acc = _mm_setzero_ps();
3012
3013    for chunk in 0..chunks {
3014        let base = chunk * 4;
3015        let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
3016            vec_u8.as_ptr().add(base) as *const i32
3017        ));
3018        let ints = _mm_cvtepu8_epi32(bytes);
3019        let floats = _mm_cvtepi32_ps(ints);
3020        let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
3021        let va = _mm_loadu_ps(query.as_ptr().add(base));
3022        acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
3023    }
3024
3025    let shuf = _mm_shuffle_ps(acc, acc, 0b10_11_00_01);
3026    let sums = _mm_add_ps(acc, shuf);
3027    let shuf2 = _mm_movehl_ps(sums, sums);
3028    let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
3029
3030    let base = chunks * 4;
3031    for i in 0..remainder {
3032        dot += *query.get_unchecked(base + i) * u8_to_f32(*vec_u8.get_unchecked(base + i));
3033    }
3034    dot
3035}
3036
3037// ============================================================================
3038// Platform dispatch
3039// ============================================================================
3040
3041#[inline]
3042fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
3043    #[cfg(target_arch = "aarch64")]
3044    {
3045        return unsafe { neon_quant::fused_dot_norm_f16(query_f16, vec_f16, dim) };
3046    }
3047
3048    #[cfg(target_arch = "x86_64")]
3049    {
3050        if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("fma") {
3051            return unsafe { fused_dot_norm_f16_f16c(query_f16, vec_f16, dim) };
3052        }
3053        if sse::is_available() {
3054            return unsafe { fused_dot_norm_f16_sse(query_f16, vec_f16, dim) };
3055        }
3056    }
3057
3058    #[allow(unreachable_code)]
3059    fused_dot_norm_f16_scalar(query_f16, vec_f16, dim)
3060}
3061
3062#[inline]
3063fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
3064    #[cfg(target_arch = "aarch64")]
3065    {
3066        return unsafe { neon_quant::fused_dot_norm_u8(query, vec_u8, dim) };
3067    }
3068
3069    #[cfg(target_arch = "x86_64")]
3070    {
3071        if sse::is_available() {
3072            return unsafe { fused_dot_norm_u8_sse(query, vec_u8, dim) };
3073        }
3074    }
3075
3076    #[allow(unreachable_code)]
3077    fused_dot_norm_u8_scalar(query, vec_u8, dim)
3078}
3079
3080// ── Dot-product-only dispatch (for unit_norm vectors) ─────────────────────
3081
3082#[inline]
3083fn dot_product_f16_quant(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
3084    #[cfg(target_arch = "aarch64")]
3085    {
3086        return unsafe { neon_quant::dot_product_f16(query_f16, vec_f16, dim) };
3087    }
3088
3089    #[cfg(target_arch = "x86_64")]
3090    {
3091        if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("fma") {
3092            return unsafe { dot_product_f16_f16c(query_f16, vec_f16, dim) };
3093        }
3094    }
3095
3096    #[allow(unreachable_code)]
3097    dot_product_f16_scalar(query_f16, vec_f16, dim)
3098}
3099
3100#[inline]
3101fn dot_product_u8_quant(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3102    #[cfg(target_arch = "aarch64")]
3103    {
3104        return unsafe { neon_quant::dot_product_u8(query, vec_u8, dim) };
3105    }
3106
3107    #[cfg(target_arch = "x86_64")]
3108    {
3109        if sse::is_available() {
3110            return unsafe { dot_product_u8_sse(query, vec_u8, dim) };
3111        }
3112    }
3113
3114    #[allow(unreachable_code)]
3115    dot_product_u8_scalar(query, vec_u8, dim)
3116}
3117
3118// ============================================================================
3119// Public batch cosine scoring for quantized vectors
3120// ============================================================================
3121
3122/// Batch cosine similarity: f32 query vs N contiguous f16 vectors.
3123///
3124/// `vectors_raw` is raw bytes: N vectors × dim × 2 bytes (f16 stored as u16).
3125/// Query is quantized to f16 once, then both query and vectors are scored in
3126/// f16 space using hardware SIMD conversion (8 elements/iteration on NEON).
3127/// Memory bandwidth is halved for both query and vector loads.
3128#[inline]
3129pub fn batch_cosine_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3130    let n = scores.len();
3131    let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3132    let required = n
3133        .checked_mul(vec_bytes)
3134        .expect("f16 batch byte length overflow");
3135    assert_eq!(
3136        query.len(),
3137        dim,
3138        "f16 batch cosine query dimension mismatch"
3139    );
3140    assert!(
3141        vectors_raw.len() >= required,
3142        "f16 batch cosine vectors are truncated: need {required} bytes, got {}",
3143        vectors_raw.len()
3144    );
3145    if required > 0 {
3146        assert!(
3147            (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3148            "f16 batch cosine vectors are not 2-byte aligned"
3149        );
3150    }
3151    if dim == 0 || n == 0 {
3152        return;
3153    }
3154
3155    // Compute query inverse norm in f32 (full precision, before quantization)
3156    let norm_q_sq = dot_product_f32(query, query, dim);
3157    if norm_q_sq < f32::EPSILON {
3158        for s in scores.iter_mut() {
3159            *s = 0.0;
3160        }
3161        return;
3162    }
3163    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3164
3165    // Quantize query to f16 once (O(dim)), reused for all N vector scorings
3166    let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3167
3168    for i in 0..n {
3169        let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3170        let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3171
3172        let (dot, norm_v_sq) = fused_dot_norm_f16(&query_f16, f16_slice, dim);
3173        scores[i] = if norm_v_sq < f32::EPSILON {
3174            0.0
3175        } else {
3176            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3177        };
3178    }
3179}
3180
3181/// Batch cosine similarity: f32 query vs N contiguous u8 vectors.
3182///
3183/// `vectors_raw` is raw bytes: N vectors × dim bytes (u8, mapping
3184/// `[-1, 1]` to `[0, 255]`).
3185/// Converts u8→f32 using NEON widening chain (16 values/iteration), scores with FMA.
3186/// Memory bandwidth is quartered compared to f32 scoring.
3187#[inline]
3188pub fn batch_cosine_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3189    let n = scores.len();
3190    let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3191    assert_eq!(query.len(), dim, "u8 batch cosine query dimension mismatch");
3192    assert!(
3193        vectors_raw.len() >= required,
3194        "u8 batch cosine vectors are truncated: need {required} bytes, got {}",
3195        vectors_raw.len()
3196    );
3197    if dim == 0 || n == 0 {
3198        return;
3199    }
3200
3201    let norm_q_sq = dot_product_f32(query, query, dim);
3202    if norm_q_sq < f32::EPSILON {
3203        for s in scores.iter_mut() {
3204            *s = 0.0;
3205        }
3206        return;
3207    }
3208    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3209
3210    for i in 0..n {
3211        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3212
3213        let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3214        scores[i] = if norm_v_sq < f32::EPSILON {
3215            0.0
3216        } else {
3217            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3218        };
3219    }
3220}
3221
3222// ============================================================================
3223// Batch dot-product scoring for unit-norm vectors
3224// ============================================================================
3225
3226/// Batch dot-product scoring: f32 query vs N contiguous f32 unit-norm vectors.
3227///
3228/// For pre-normalized vectors (||v|| = 1), cosine = dot(q, v) / ||q||.
3229/// Skips per-vector norm computation — ~40% less work than `batch_cosine_scores`.
3230#[inline]
3231pub fn batch_dot_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
3232    let n = scores.len();
3233    let required = n
3234        .checked_mul(dim)
3235        .expect("batch dot vector length overflow");
3236    assert_eq!(query.len(), dim, "batch dot query dimension mismatch");
3237    assert!(
3238        vectors.len() >= required,
3239        "batch dot vectors are truncated: need {required}, got {}",
3240        vectors.len()
3241    );
3242
3243    if dim == 0 || n == 0 {
3244        return;
3245    }
3246
3247    let norm_q_sq = dot_product_f32(query, query, dim);
3248    if norm_q_sq < f32::EPSILON {
3249        for s in scores.iter_mut() {
3250            *s = 0.0;
3251        }
3252        return;
3253    }
3254    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3255
3256    for i in 0..n {
3257        let vec = &vectors[i * dim..(i + 1) * dim];
3258        let dot = dot_product_f32(query, vec, dim);
3259        scores[i] = dot * inv_norm_q;
3260    }
3261}
3262
3263/// Batch dot-product scoring: f32 query vs N contiguous f16 unit-norm vectors.
3264///
3265/// For pre-normalized vectors (||v|| = 1), cosine = dot(q, v) / ||q||.
3266/// Uses F16C/NEON hardware conversion + dot-only kernel.
3267#[inline]
3268pub fn batch_dot_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3269    let n = scores.len();
3270    let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3271    let required = n
3272        .checked_mul(vec_bytes)
3273        .expect("f16 batch byte length overflow");
3274    assert_eq!(query.len(), dim, "f16 batch dot query dimension mismatch");
3275    assert!(
3276        vectors_raw.len() >= required,
3277        "f16 batch dot vectors are truncated: need {required} bytes, got {}",
3278        vectors_raw.len()
3279    );
3280    if required > 0 {
3281        assert!(
3282            (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3283            "f16 batch dot vectors are not 2-byte aligned"
3284        );
3285    }
3286    if dim == 0 || n == 0 {
3287        return;
3288    }
3289
3290    let norm_q_sq = dot_product_f32(query, query, dim);
3291    if norm_q_sq < f32::EPSILON {
3292        for s in scores.iter_mut() {
3293            *s = 0.0;
3294        }
3295        return;
3296    }
3297    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3298
3299    let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3300    for i in 0..n {
3301        let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3302        let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3303        let dot = dot_product_f16_quant(&query_f16, f16_slice, dim);
3304        scores[i] = dot * inv_norm_q;
3305    }
3306}
3307
3308/// Batch dot-product scoring: f32 query vs N contiguous u8 unit-norm vectors.
3309///
3310/// For pre-normalized vectors (||v|| = 1), cosine = dot(q, v) / ||q||.
3311/// Uses NEON/SSE widening chain for u8→f32 conversion + dot-only kernel.
3312#[inline]
3313pub fn batch_dot_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3314    let n = scores.len();
3315    let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3316    assert_eq!(query.len(), dim, "u8 batch dot query dimension mismatch");
3317    assert!(
3318        vectors_raw.len() >= required,
3319        "u8 batch dot vectors are truncated: need {required} bytes, got {}",
3320        vectors_raw.len()
3321    );
3322    if dim == 0 || n == 0 {
3323        return;
3324    }
3325
3326    let norm_q_sq = dot_product_f32(query, query, dim);
3327    if norm_q_sq < f32::EPSILON {
3328        for s in scores.iter_mut() {
3329            *s = 0.0;
3330        }
3331        return;
3332    }
3333    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3334
3335    for i in 0..n {
3336        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3337        let dot = dot_product_u8_quant(query, u8_slice, dim);
3338        scores[i] = dot * inv_norm_q;
3339    }
3340}
3341
3342// ============================================================================
3343// Precomputed-norm batch scoring (avoids redundant query norm + f16 conversion)
3344// ============================================================================
3345
3346/// Batch cosine: f32 query vs N f32 vectors, with precomputed `inv_norm_q`.
3347#[inline]
3348pub fn batch_cosine_scores_precomp(
3349    query: &[f32],
3350    vectors: &[f32],
3351    dim: usize,
3352    scores: &mut [f32],
3353    inv_norm_q: f32,
3354) {
3355    let n = scores.len();
3356    let required = n
3357        .checked_mul(dim)
3358        .expect("precomputed cosine vector length overflow");
3359    assert_eq!(
3360        query.len(),
3361        dim,
3362        "precomputed cosine query dimension mismatch"
3363    );
3364    assert!(
3365        vectors.len() >= required,
3366        "precomputed cosine vectors are truncated: need {required}, got {}",
3367        vectors.len()
3368    );
3369    for i in 0..n {
3370        let vec = &vectors[i * dim..(i + 1) * dim];
3371        let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
3372        scores[i] = if norm_v_sq < f32::EPSILON {
3373            0.0
3374        } else {
3375            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3376        };
3377    }
3378}
3379
3380/// Batch cosine: precomputed `inv_norm_q` + `query_f16` vs N f16 vectors.
3381#[inline]
3382pub fn batch_cosine_scores_f16_precomp(
3383    query_f16: &[u16],
3384    vectors_raw: &[u8],
3385    dim: usize,
3386    scores: &mut [f32],
3387    inv_norm_q: f32,
3388) {
3389    let n = scores.len();
3390    let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3391    let required = n
3392        .checked_mul(vec_bytes)
3393        .expect("precomputed f16 cosine batch byte length overflow");
3394    assert_eq!(
3395        query_f16.len(),
3396        dim,
3397        "precomputed f16 cosine query dimension mismatch"
3398    );
3399    assert!(
3400        vectors_raw.len() >= required,
3401        "precomputed f16 cosine vectors are truncated: need {required} bytes, got {}",
3402        vectors_raw.len()
3403    );
3404    if required > 0 {
3405        assert!(
3406            (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3407            "precomputed f16 cosine vectors are not 2-byte aligned"
3408        );
3409    }
3410    for i in 0..n {
3411        let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3412        let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3413        let (dot, norm_v_sq) = fused_dot_norm_f16(query_f16, f16_slice, dim);
3414        scores[i] = if norm_v_sq < f32::EPSILON {
3415            0.0
3416        } else {
3417            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3418        };
3419    }
3420}
3421
3422/// Batch cosine: precomputed `inv_norm_q` vs N u8 vectors.
3423#[inline]
3424pub fn batch_cosine_scores_u8_precomp(
3425    query: &[f32],
3426    vectors_raw: &[u8],
3427    dim: usize,
3428    scores: &mut [f32],
3429    inv_norm_q: f32,
3430) {
3431    let n = scores.len();
3432    let required = n
3433        .checked_mul(dim)
3434        .expect("precomputed u8 cosine batch byte length overflow");
3435    assert_eq!(
3436        query.len(),
3437        dim,
3438        "precomputed u8 cosine query dimension mismatch"
3439    );
3440    assert!(
3441        vectors_raw.len() >= required,
3442        "precomputed u8 cosine vectors are truncated: need {required} bytes, got {}",
3443        vectors_raw.len()
3444    );
3445    for i in 0..n {
3446        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3447        let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3448        scores[i] = if norm_v_sq < f32::EPSILON {
3449            0.0
3450        } else {
3451            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3452        };
3453    }
3454}
3455
3456/// Batch dot-product: precomputed `inv_norm_q` vs N f32 unit-norm vectors.
3457#[inline]
3458pub fn batch_dot_scores_precomp(
3459    query: &[f32],
3460    vectors: &[f32],
3461    dim: usize,
3462    scores: &mut [f32],
3463    inv_norm_q: f32,
3464) {
3465    let n = scores.len();
3466    let required = n
3467        .checked_mul(dim)
3468        .expect("precomputed dot vector length overflow");
3469    assert_eq!(query.len(), dim, "precomputed dot query dimension mismatch");
3470    assert!(
3471        vectors.len() >= required,
3472        "precomputed dot vectors are truncated: need {required}, got {}",
3473        vectors.len()
3474    );
3475    for i in 0..n {
3476        let vec = &vectors[i * dim..(i + 1) * dim];
3477        scores[i] = dot_product_f32(query, vec, dim) * inv_norm_q;
3478    }
3479}
3480
3481/// Batch dot-product: precomputed `inv_norm_q` + `query_f16` vs N f16 unit-norm vectors.
3482#[inline]
3483pub fn batch_dot_scores_f16_precomp(
3484    query_f16: &[u16],
3485    vectors_raw: &[u8],
3486    dim: usize,
3487    scores: &mut [f32],
3488    inv_norm_q: f32,
3489) {
3490    let n = scores.len();
3491    let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3492    let required = n
3493        .checked_mul(vec_bytes)
3494        .expect("precomputed f16 dot batch byte length overflow");
3495    assert_eq!(
3496        query_f16.len(),
3497        dim,
3498        "precomputed f16 dot query dimension mismatch"
3499    );
3500    assert!(
3501        vectors_raw.len() >= required,
3502        "precomputed f16 dot vectors are truncated: need {required} bytes, got {}",
3503        vectors_raw.len()
3504    );
3505    if required > 0 {
3506        assert!(
3507            (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3508            "precomputed f16 dot vectors are not 2-byte aligned"
3509        );
3510    }
3511    for i in 0..n {
3512        let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3513        let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3514        scores[i] = dot_product_f16_quant(query_f16, f16_slice, dim) * inv_norm_q;
3515    }
3516}
3517
3518/// Batch dot-product: precomputed `inv_norm_q` vs N u8 unit-norm vectors.
3519#[inline]
3520pub fn batch_dot_scores_u8_precomp(
3521    query: &[f32],
3522    vectors_raw: &[u8],
3523    dim: usize,
3524    scores: &mut [f32],
3525    inv_norm_q: f32,
3526) {
3527    let n = scores.len();
3528    let required = n
3529        .checked_mul(dim)
3530        .expect("precomputed u8 dot batch byte length overflow");
3531    assert_eq!(
3532        query.len(),
3533        dim,
3534        "precomputed u8 dot query dimension mismatch"
3535    );
3536    assert!(
3537        vectors_raw.len() >= required,
3538        "precomputed u8 dot vectors are truncated: need {required} bytes, got {}",
3539        vectors_raw.len()
3540    );
3541    for i in 0..n {
3542        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3543        scores[i] = dot_product_u8_quant(query, u8_slice, dim) * inv_norm_q;
3544    }
3545}
3546
3547/// Compute cosine similarity between two f32 vectors with SIMD acceleration
3548///
3549/// Returns dot(a,b) / (||a|| * ||b||), range [-1, 1]
3550/// Returns 0.0 if either vector has zero norm.
3551#[inline]
3552pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
3553    assert_eq!(a.len(), b.len(), "cosine vector dimension mismatch");
3554    let count = a.len();
3555
3556    if count == 0 {
3557        return 0.0;
3558    }
3559
3560    let dot = dot_product_f32(a, b, count);
3561    let norm_a = dot_product_f32(a, a, count);
3562    let norm_b = dot_product_f32(b, b, count);
3563
3564    let denom = (norm_a * norm_b).sqrt();
3565    if denom < f32::EPSILON {
3566        return 0.0;
3567    }
3568
3569    dot / denom
3570}
3571
3572// ============================================================================
3573// Hamming distance for binary dense vectors
3574// ============================================================================
3575
3576/// AVX-512 Hamming distance using `VPOPCNTDQ`.
3577///
3578/// Processes 64 bytes per iteration with a single hardware popcount per lane
3579/// group, which removes the nibble-lookup shuffles the AVX2 path needs.
3580#[cfg(target_arch = "x86_64")]
3581#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3582#[allow(unsafe_op_in_unsafe_fn)]
3583unsafe fn hamming_distance_avx512(a: &[u8], b: &[u8]) -> u32 {
3584    use std::arch::x86_64::*;
3585
3586    let len = a.len();
3587    let chunks64 = len / 64;
3588    let mut acc = _mm512_setzero_si512();
3589
3590    for c in 0..chunks64 {
3591        let off = c * 64;
3592        let va = _mm512_loadu_si512(a.as_ptr().add(off) as *const __m512i);
3593        let vb = _mm512_loadu_si512(b.as_ptr().add(off) as *const __m512i);
3594        acc = _mm512_add_epi64(acc, _mm512_popcnt_epi64(_mm512_xor_si512(va, vb)));
3595    }
3596
3597    let base = chunks64 * 64;
3598    _mm512_reduce_add_epi64(acc) as u32 + hamming_distance_scalar(&a[base..], &b[base..])
3599}
3600
3601/// Four-row AVX-512 Hamming distance sharing the query load across rows.
3602#[cfg(target_arch = "x86_64")]
3603#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3604#[allow(unsafe_op_in_unsafe_fn)]
3605unsafe fn hamming_distance_x4_avx512(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3606    use std::arch::x86_64::*;
3607
3608    let len = query.len();
3609    let chunks64 = len / 64;
3610    let mut acc = [_mm512_setzero_si512(); 4];
3611
3612    for c in 0..chunks64 {
3613        let off = c * 64;
3614        let vq = _mm512_loadu_si512(query.as_ptr().add(off) as *const __m512i);
3615        for r in 0..4 {
3616            let vr = _mm512_loadu_si512(rows[r].as_ptr().add(off) as *const __m512i);
3617            acc[r] = _mm512_add_epi64(acc[r], _mm512_popcnt_epi64(_mm512_xor_si512(vq, vr)));
3618        }
3619    }
3620
3621    let base = chunks64 * 64;
3622    let tail = &query[base..];
3623    [
3624        _mm512_reduce_add_epi64(acc[0]) as u32 + hamming_distance_scalar(tail, &rows[0][base..]),
3625        _mm512_reduce_add_epi64(acc[1]) as u32 + hamming_distance_scalar(tail, &rows[1][base..]),
3626        _mm512_reduce_add_epi64(acc[2]) as u32 + hamming_distance_scalar(tail, &rows[2][base..]),
3627        _mm512_reduce_add_epi64(acc[3]) as u32 + hamming_distance_scalar(tail, &rows[3][base..]),
3628    ]
3629}
3630
3631/// Four-row scalar Hamming distance sharing the query load across rows.
3632#[inline]
3633fn hamming_distance_x4_scalar(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3634    let len = query.len();
3635    let chunks = len / 8;
3636    let mut total = [0u32; 4];
3637
3638    for i in 0..chunks {
3639        let off = i * 8;
3640        let vq = unsafe { std::ptr::read_unaligned(query.as_ptr().add(off) as *const u64) };
3641        for r in 0..4 {
3642            let vr = unsafe { std::ptr::read_unaligned(rows[r].as_ptr().add(off) as *const u64) };
3643            total[r] += (vq ^ vr).count_ones();
3644        }
3645    }
3646
3647    let base = chunks * 8;
3648    for k in base..len {
3649        let q = query[k];
3650        for r in 0..4 {
3651            total[r] += (q ^ rows[r][k]).count_ones();
3652        }
3653    }
3654
3655    total
3656}
3657
3658/// Rows scored per kernel invocation. Sharing the query load, the AVX2 nibble
3659/// lookup table and the horizontal reduction across four rows amortises the
3660/// non-inlinable `#[target_feature]` call and overlaps the popcount chains.
3661const HAMMING_ROWS_PER_KERNEL: usize = 4;
3662
3663/// Architecture kernel resolved once for a whole scan.
3664///
3665/// Hot binary paths — HNSW centroid routing, k-majority assignment, leaf
3666/// scanning — score millions of code pairs against one query. Resolving the
3667/// kernel up front keeps runtime feature detection out of the inner loop, and
3668/// the row-batched entry points let one dispatch cover a whole neighbour list.
3669#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3670pub enum HammingKernel {
3671    #[cfg(target_arch = "x86_64")]
3672    Avx512,
3673    #[cfg(target_arch = "x86_64")]
3674    Avx2,
3675    #[cfg(target_arch = "aarch64")]
3676    Neon,
3677    Scalar,
3678}
3679
3680impl HammingKernel {
3681    /// Detect the widest kernel this CPU supports.
3682    #[inline]
3683    pub fn resolve() -> Self {
3684        #[cfg(target_arch = "x86_64")]
3685        {
3686            if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512vpopcntdq") {
3687                return Self::Avx512;
3688            }
3689            if avx2::is_available() {
3690                return Self::Avx2;
3691            }
3692            Self::Scalar
3693        }
3694
3695        #[cfg(target_arch = "aarch64")]
3696        {
3697            Self::Neon
3698        }
3699
3700        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
3701        {
3702            Self::Scalar
3703        }
3704    }
3705
3706    /// Hamming distance between two equal-length packed-bit vectors.
3707    #[inline]
3708    pub fn distance(self, a: &[u8], b: &[u8]) -> u32 {
3709        debug_assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3710        match self {
3711            #[cfg(target_arch = "x86_64")]
3712            Self::Avx512 => unsafe { hamming_distance_avx512(a, b) },
3713            #[cfg(target_arch = "x86_64")]
3714            Self::Avx2 => unsafe { avx2::hamming_distance(a, b) },
3715            #[cfg(target_arch = "aarch64")]
3716            Self::Neon => unsafe { neon::hamming_distance(a, b) },
3717            Self::Scalar => hamming_distance_scalar(a, b),
3718        }
3719    }
3720
3721    /// `out[i]` receives the distance from `query` to row `i` of `db`.
3722    pub fn distances(self, query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3723        self.score_rows(query, db, byte_len, out, |index| index);
3724    }
3725
3726    /// `out[i]` receives the distance from `query` to row `ids[i]` of `db`.
3727    ///
3728    /// Graph routing visits scattered centroid rows; gathering them through one
3729    /// dispatch keeps the batched kernel usable there.
3730    pub fn gather_distances(
3731        self,
3732        query: &[u8],
3733        db: &[u8],
3734        byte_len: usize,
3735        ids: &[u32],
3736        out: &mut [u32],
3737    ) {
3738        assert_eq!(
3739            ids.len(),
3740            out.len(),
3741            "Hamming gather needs one output slot per row id"
3742        );
3743        self.score_rows(query, db, byte_len, out, |index| ids[index] as usize);
3744    }
3745
3746    #[inline]
3747    fn score_rows(
3748        self,
3749        query: &[u8],
3750        db: &[u8],
3751        byte_len: usize,
3752        out: &mut [u32],
3753        index_of: impl Fn(usize) -> usize,
3754    ) {
3755        assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3756        if byte_len == 0 || out.is_empty() {
3757            return;
3758        }
3759        let row = |index: usize| -> &[u8] {
3760            let start = index * byte_len;
3761            &db[start..start + byte_len]
3762        };
3763        macro_rules! score_with {
3764            ($one:expr, $four:expr) => {{
3765                let mut i = 0;
3766                while i + HAMMING_ROWS_PER_KERNEL <= out.len() {
3767                    let quad = [
3768                        row(index_of(i)),
3769                        row(index_of(i + 1)),
3770                        row(index_of(i + 2)),
3771                        row(index_of(i + 3)),
3772                    ];
3773                    out[i..i + HAMMING_ROWS_PER_KERNEL].copy_from_slice(&$four(query, quad));
3774                    i += HAMMING_ROWS_PER_KERNEL;
3775                }
3776                while i < out.len() {
3777                    out[i] = $one(query, row(index_of(i)));
3778                    i += 1;
3779                }
3780            }};
3781        }
3782        match self {
3783            #[cfg(target_arch = "x86_64")]
3784            Self::Avx512 => score_with!(
3785                |query, row| unsafe { hamming_distance_avx512(query, row) },
3786                |query, rows| unsafe { hamming_distance_x4_avx512(query, rows) }
3787            ),
3788            #[cfg(target_arch = "x86_64")]
3789            Self::Avx2 => score_with!(
3790                |query, row| unsafe { avx2::hamming_distance(query, row) },
3791                |query, rows| unsafe { avx2::hamming_distance_x4(query, rows) }
3792            ),
3793            #[cfg(target_arch = "aarch64")]
3794            Self::Neon => score_with!(
3795                |query, row| unsafe { neon::hamming_distance(query, row) },
3796                |query, rows| unsafe { neon::hamming_distance_x4(query, rows) }
3797            ),
3798            Self::Scalar => score_with!(hamming_distance_scalar, hamming_distance_x4_scalar),
3799        }
3800    }
3801}
3802
3803/// Compute Hamming distance between two packed-bit vectors.
3804/// Returns the number of differing bits.
3805///
3806/// Uses NEON on aarch64 and VPOPCNTDQ/AVX2 on x86_64, with a scalar fallback.
3807/// Loops over many pairs should resolve a [`HammingKernel`] once instead of
3808/// paying feature detection here per pair.
3809#[inline]
3810pub fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
3811    assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3812    HammingKernel::resolve().distance(a, b)
3813}
3814
3815/// Scalar Hamming distance using u64 chunks + count_ones().
3816/// On x86_64, count_ones() compiles to POPCNT when target-cpu supports it.
3817#[inline]
3818#[allow(dead_code)]
3819fn hamming_distance_scalar(a: &[u8], b: &[u8]) -> u32 {
3820    let len = a.len();
3821    let chunks = len / 8;
3822    let remainder = len % 8;
3823    let mut total = 0u32;
3824
3825    for i in 0..chunks {
3826        let off = i * 8;
3827        let va = unsafe { std::ptr::read_unaligned(a.as_ptr().add(off) as *const u64) };
3828        let vb = unsafe { std::ptr::read_unaligned(b.as_ptr().add(off) as *const u64) };
3829        total += (va ^ vb).count_ones();
3830    }
3831
3832    let base = chunks * 8;
3833    for i in 0..remainder {
3834        total += (a[base + i] ^ b[base + i]).count_ones();
3835    }
3836
3837    total
3838}
3839
3840/// Batch Hamming scoring: compute similarity scores for multiple binary vectors.
3841///
3842/// `query` and each vector in `db` are packed-bit vectors of `byte_len` bytes each.
3843/// `dim_bits` is the number of bits (dimensions) for normalization.
3844/// Score = 1.0 - hamming_distance / dim_bits (range [0.0, 1.0]).
3845pub fn batch_hamming_scores(
3846    query: &[u8],
3847    db: &[u8],
3848    byte_len: usize,
3849    dim_bits: usize,
3850    scores: &mut [f32],
3851) {
3852    let n = scores.len();
3853    let required = n
3854        .checked_mul(byte_len)
3855        .expect("Hamming batch byte length overflow");
3856    assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3857    assert!(
3858        db.len() >= required,
3859        "Hamming batch is truncated: need {required} bytes, got {}",
3860        db.len()
3861    );
3862
3863    if byte_len == 0 || n == 0 || dim_bits == 0 {
3864        return;
3865    }
3866
3867    scores_from_hamming(
3868        HammingKernel::resolve(),
3869        query,
3870        db,
3871        byte_len,
3872        dim_bits,
3873        scores,
3874    );
3875}
3876
3877/// Batch Hamming scoring with a caller-resolved kernel.
3878///
3879/// Scans that already hold a [`HammingKernel`] (leaf scanning, Lloyd
3880/// assignment) use this to keep feature detection out of the loop entirely.
3881pub fn scores_from_hamming(
3882    kernel: HammingKernel,
3883    query: &[u8],
3884    db: &[u8],
3885    byte_len: usize,
3886    dim_bits: usize,
3887    scores: &mut [f32],
3888) {
3889    if byte_len == 0 || scores.is_empty() || dim_bits == 0 {
3890        return;
3891    }
3892    let inv_dim = 1.0 / dim_bits as f32;
3893    // Distances stay integral until the very last step; the stack block keeps
3894    // the row-batched kernel reachable without a per-scan allocation.
3895    let mut distances = [0u32; HAMMING_DISTANCE_BLOCK];
3896    for (block_index, block) in scores.chunks_mut(HAMMING_DISTANCE_BLOCK).enumerate() {
3897        let rows = &mut distances[..block.len()];
3898        kernel.distances(
3899            query,
3900            &db[block_index * HAMMING_DISTANCE_BLOCK * byte_len..],
3901            byte_len,
3902            rows,
3903        );
3904        for (score, &distance) in block.iter_mut().zip(rows.iter()) {
3905            *score = 1.0 - distance as f32 * inv_dim;
3906        }
3907    }
3908}
3909
3910/// Rows per stack block when converting batched distances into scores.
3911const HAMMING_DISTANCE_BLOCK: usize = 64;
3912
3913/// Batch Hamming distances (exact bit counts) for `out.len()` rows of `db`.
3914///
3915/// Callers that rank by distance — coarse assignment, routing — avoid the
3916/// float round-trip entirely.
3917pub fn batch_hamming_distances(query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3918    HammingKernel::resolve().distances(query, db, byte_len, out);
3919}
3920
3921#[cfg(test)]
3922mod tests {
3923    use super::*;
3924
3925    #[test]
3926    fn vector_simd_boundaries_reject_dimension_mismatches() {
3927        let vectors = vec![1.0f32; 6];
3928        let raw_f16 = vec![0u8; 12];
3929        let raw_u8 = vec![0u8; 6];
3930        let mut scores = vec![0.0f32; 2];
3931
3932        for invalid_query in [vec![1.0, 2.0], vec![1.0, 2.0, 3.0, 4.0]] {
3933            assert!(
3934                std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3935                    batch_cosine_scores(&invalid_query, &vectors, 3, &mut scores)
3936                }))
3937                .is_err()
3938            );
3939            assert!(
3940                std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3941                    batch_dot_scores_f16(&invalid_query, &raw_f16, 3, &mut scores)
3942                }))
3943                .is_err()
3944            );
3945            assert!(
3946                std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3947                    batch_cosine_scores_u8(&invalid_query, &raw_u8, 3, &mut scores)
3948                }))
3949                .is_err()
3950            );
3951        }
3952    }
3953
3954    #[test]
3955    fn vector_simd_boundaries_reject_truncated_storage() {
3956        let query = [1.0f32, 2.0, 3.0];
3957        let mut scores = [0.0f32; 2];
3958
3959        assert!(
3960            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3961                batch_dot_scores(&query, &[0.0; 5], 3, &mut scores)
3962            }))
3963            .is_err()
3964        );
3965        assert!(
3966            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3967                batch_cosine_scores_f16(&query, &[0u8; 11], 3, &mut scores)
3968            }))
3969            .is_err()
3970        );
3971        assert!(
3972            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3973                dot_product_f32(&query, &query, 4)
3974            }))
3975            .is_err()
3976        );
3977    }
3978
3979    #[test]
3980    fn test_unpack_8bit() {
3981        let input: Vec<u8> = (0..128).collect();
3982        let mut output = vec![0u32; 128];
3983        unpack_8bit(&input, &mut output, 128);
3984
3985        for (i, &v) in output.iter().enumerate() {
3986            assert_eq!(v, i as u32);
3987        }
3988    }
3989
3990    #[test]
3991    fn test_unpack_16bit() {
3992        let mut input = vec![0u8; 256];
3993        for i in 0..128 {
3994            let val = (i * 100) as u16;
3995            input[i * 2] = val as u8;
3996            input[i * 2 + 1] = (val >> 8) as u8;
3997        }
3998
3999        let mut output = vec![0u32; 128];
4000        unpack_16bit(&input, &mut output, 128);
4001
4002        for (i, &v) in output.iter().enumerate() {
4003            assert_eq!(v, (i * 100) as u32);
4004        }
4005    }
4006
4007    #[test]
4008    fn test_unpack_32bit() {
4009        let mut input = vec![0u8; 512];
4010        for i in 0..128 {
4011            let val = (i * 1000) as u32;
4012            let bytes = val.to_le_bytes();
4013            input[i * 4..i * 4 + 4].copy_from_slice(&bytes);
4014        }
4015
4016        let mut output = vec![0u32; 128];
4017        unpack_32bit(&input, &mut output, 128);
4018
4019        for (i, &v) in output.iter().enumerate() {
4020            assert_eq!(v, (i * 1000) as u32);
4021        }
4022    }
4023
4024    #[test]
4025    fn test_delta_decode() {
4026        // doc_ids: [10, 15, 20, 30, 50]
4027        // gaps: [5, 5, 10, 20]
4028        // deltas (gap-1): [4, 4, 9, 19]
4029        let deltas = vec![4u32, 4, 9, 19];
4030        let mut output = vec![0u32; 5];
4031
4032        delta_decode(&mut output, &deltas, 10, 5);
4033
4034        assert_eq!(output, vec![10, 15, 20, 30, 50]);
4035    }
4036
4037    #[test]
4038    fn test_add_one() {
4039        let mut values = vec![0u32, 1, 2, 3, 4, 5, 6, 7];
4040        add_one(&mut values, 8);
4041
4042        assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
4043    }
4044
4045    #[test]
4046    fn test_bits_needed() {
4047        assert_eq!(bits_needed(0), 0);
4048        assert_eq!(bits_needed(1), 1);
4049        assert_eq!(bits_needed(2), 2);
4050        assert_eq!(bits_needed(3), 2);
4051        assert_eq!(bits_needed(4), 3);
4052        assert_eq!(bits_needed(255), 8);
4053        assert_eq!(bits_needed(256), 9);
4054        assert_eq!(bits_needed(u32::MAX), 32);
4055    }
4056
4057    #[test]
4058    fn test_unpack_8bit_delta_decode() {
4059        // doc_ids: [10, 15, 20, 30, 50]
4060        // gaps: [5, 5, 10, 20]
4061        // deltas (gap-1): [4, 4, 9, 19] stored as u8
4062        let input: Vec<u8> = vec![4, 4, 9, 19];
4063        let mut output = vec![0u32; 5];
4064
4065        unpack_8bit_delta_decode(&input, &mut output, 10, 5);
4066
4067        assert_eq!(output, vec![10, 15, 20, 30, 50]);
4068    }
4069
4070    #[test]
4071    fn test_unpack_16bit_delta_decode() {
4072        // doc_ids: [100, 600, 1100, 2100, 4100]
4073        // gaps: [500, 500, 1000, 2000]
4074        // deltas (gap-1): [499, 499, 999, 1999] stored as u16
4075        let mut input = vec![0u8; 8];
4076        for (i, &delta) in [499u16, 499, 999, 1999].iter().enumerate() {
4077            input[i * 2] = delta as u8;
4078            input[i * 2 + 1] = (delta >> 8) as u8;
4079        }
4080        let mut output = vec![0u32; 5];
4081
4082        unpack_16bit_delta_decode(&input, &mut output, 100, 5);
4083
4084        assert_eq!(output, vec![100, 600, 1100, 2100, 4100]);
4085    }
4086
4087    #[test]
4088    fn test_fused_vs_separate_8bit() {
4089        // Test that fused and separate operations produce the same result
4090        let input: Vec<u8> = (0..127).collect();
4091        let first_value = 1000u32;
4092        let count = 128;
4093
4094        // Separate: unpack then delta_decode
4095        let mut unpacked = vec![0u32; 128];
4096        unpack_8bit(&input, &mut unpacked, 127);
4097        let mut separate_output = vec![0u32; 128];
4098        delta_decode(&mut separate_output, &unpacked, first_value, count);
4099
4100        // Fused
4101        let mut fused_output = vec![0u32; 128];
4102        unpack_8bit_delta_decode(&input, &mut fused_output, first_value, count);
4103
4104        assert_eq!(separate_output, fused_output);
4105    }
4106
4107    #[test]
4108    fn test_round_bit_width() {
4109        assert_eq!(round_bit_width(0), 0);
4110        assert_eq!(round_bit_width(1), 8);
4111        assert_eq!(round_bit_width(5), 8);
4112        assert_eq!(round_bit_width(8), 8);
4113        assert_eq!(round_bit_width(9), 16);
4114        assert_eq!(round_bit_width(12), 16);
4115        assert_eq!(round_bit_width(16), 16);
4116        assert_eq!(round_bit_width(17), 32);
4117        assert_eq!(round_bit_width(24), 32);
4118        assert_eq!(round_bit_width(32), 32);
4119    }
4120
4121    #[test]
4122    fn test_rounded_bitwidth_from_exact() {
4123        assert_eq!(RoundedBitWidth::from_exact(0), RoundedBitWidth::Zero);
4124        assert_eq!(RoundedBitWidth::from_exact(1), RoundedBitWidth::Bits8);
4125        assert_eq!(RoundedBitWidth::from_exact(8), RoundedBitWidth::Bits8);
4126        assert_eq!(RoundedBitWidth::from_exact(9), RoundedBitWidth::Bits16);
4127        assert_eq!(RoundedBitWidth::from_exact(16), RoundedBitWidth::Bits16);
4128        assert_eq!(RoundedBitWidth::from_exact(17), RoundedBitWidth::Bits32);
4129        assert_eq!(RoundedBitWidth::from_exact(32), RoundedBitWidth::Bits32);
4130    }
4131
4132    #[test]
4133    fn test_pack_unpack_rounded_8bit() {
4134        let values: Vec<u32> = (0..128).map(|i| i % 256).collect();
4135        let mut packed = vec![0u8; 128];
4136
4137        let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits8, &mut packed);
4138        assert_eq!(bytes_written, 128);
4139
4140        let mut unpacked = vec![0u32; 128];
4141        unpack_rounded(&packed, RoundedBitWidth::Bits8, &mut unpacked, 128);
4142
4143        assert_eq!(values, unpacked);
4144    }
4145
4146    #[test]
4147    fn test_pack_unpack_rounded_16bit() {
4148        let values: Vec<u32> = (0..128).map(|i| i * 100).collect();
4149        let mut packed = vec![0u8; 256];
4150
4151        let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits16, &mut packed);
4152        assert_eq!(bytes_written, 256);
4153
4154        let mut unpacked = vec![0u32; 128];
4155        unpack_rounded(&packed, RoundedBitWidth::Bits16, &mut unpacked, 128);
4156
4157        assert_eq!(values, unpacked);
4158    }
4159
4160    #[test]
4161    fn test_pack_unpack_rounded_32bit() {
4162        let values: Vec<u32> = (0..128).map(|i| i * 100000).collect();
4163        let mut packed = vec![0u8; 512];
4164
4165        let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits32, &mut packed);
4166        assert_eq!(bytes_written, 512);
4167
4168        let mut unpacked = vec![0u32; 128];
4169        unpack_rounded(&packed, RoundedBitWidth::Bits32, &mut unpacked, 128);
4170
4171        assert_eq!(values, unpacked);
4172    }
4173
4174    #[test]
4175    fn test_unpack_rounded_delta_decode() {
4176        // Test 8-bit rounded delta decode
4177        // doc_ids: [10, 15, 20, 30, 50]
4178        // gaps: [5, 5, 10, 20]
4179        // deltas (gap-1): [4, 4, 9, 19] stored as u8
4180        let input: Vec<u8> = vec![4, 4, 9, 19];
4181        let mut output = vec![0u32; 5];
4182
4183        unpack_rounded_delta_decode(&input, RoundedBitWidth::Bits8, &mut output, 10, 5);
4184
4185        assert_eq!(output, vec![10, 15, 20, 30, 50]);
4186    }
4187
4188    #[test]
4189    fn test_unpack_rounded_delta_decode_zero() {
4190        // All zeros means gaps of 1 (consecutive doc IDs)
4191        let input: Vec<u8> = vec![];
4192        let mut output = vec![0u32; 5];
4193
4194        unpack_rounded_delta_decode(&input, RoundedBitWidth::Zero, &mut output, 100, 5);
4195
4196        assert_eq!(output, vec![100, 101, 102, 103, 104]);
4197    }
4198
4199    // ========================================================================
4200    // Sparse Vector SIMD Tests
4201    // ========================================================================
4202
4203    #[test]
4204    fn test_dequantize_uint8() {
4205        let input: Vec<u8> = vec![0, 128, 255, 64, 192];
4206        let mut output = vec![0.0f32; 5];
4207        let scale = 0.1;
4208        let min_val = 1.0;
4209
4210        dequantize_uint8(&input, &mut output, scale, min_val, 5);
4211
4212        // Expected: input[i] * scale + min_val
4213        assert!((output[0] - 1.0).abs() < 1e-6); // 0 * 0.1 + 1.0 = 1.0
4214        assert!((output[1] - 13.8).abs() < 1e-6); // 128 * 0.1 + 1.0 = 13.8
4215        assert!((output[2] - 26.5).abs() < 1e-6); // 255 * 0.1 + 1.0 = 26.5
4216        assert!((output[3] - 7.4).abs() < 1e-6); // 64 * 0.1 + 1.0 = 7.4
4217        assert!((output[4] - 20.2).abs() < 1e-6); // 192 * 0.1 + 1.0 = 20.2
4218    }
4219
4220    #[test]
4221    fn test_dequantize_uint8_large() {
4222        // Test with 128 values (full SIMD block)
4223        let input: Vec<u8> = (0..128).collect();
4224        let mut output = vec![0.0f32; 128];
4225        let scale = 2.0;
4226        let min_val = -10.0;
4227
4228        dequantize_uint8(&input, &mut output, scale, min_val, 128);
4229
4230        for (i, &out) in output.iter().enumerate().take(128) {
4231            let expected = i as f32 * scale + min_val;
4232            assert!(
4233                (out - expected).abs() < 1e-5,
4234                "Mismatch at {}: expected {}, got {}",
4235                i,
4236                expected,
4237                out
4238            );
4239        }
4240    }
4241
4242    #[test]
4243    fn test_dot_product_f32() {
4244        let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
4245        let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0];
4246
4247        let result = dot_product_f32(&a, &b, 5);
4248
4249        // Expected: 1*2 + 2*3 + 3*4 + 4*5 + 5*6 = 2 + 6 + 12 + 20 + 30 = 70
4250        assert!((result - 70.0).abs() < 1e-5);
4251    }
4252
4253    #[test]
4254    fn test_dot_product_f32_large() {
4255        // Test with 128 values
4256        let a: Vec<f32> = (0..128).map(|i| i as f32).collect();
4257        let b: Vec<f32> = (0..128).map(|i| (i + 1) as f32).collect();
4258
4259        let result = dot_product_f32(&a, &b, 128);
4260
4261        // Compute expected
4262        let expected: f32 = (0..128).map(|i| (i as f32) * ((i + 1) as f32)).sum();
4263        assert!(
4264            (result - expected).abs() < 1e-3,
4265            "Expected {}, got {}",
4266            expected,
4267            result
4268        );
4269    }
4270
4271    #[test]
4272    fn test_fused_dot_norm() {
4273        let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
4274        let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
4275        let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4276
4277        let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4278        let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4279        assert!(
4280            (dot - expected_dot).abs() < 1e-5,
4281            "dot: expected {}, got {}",
4282            expected_dot,
4283            dot
4284        );
4285        assert!(
4286            (norm_b - expected_norm).abs() < 1e-5,
4287            "norm: expected {}, got {}",
4288            expected_norm,
4289            norm_b
4290        );
4291    }
4292
4293    #[test]
4294    fn test_fused_dot_norm_large() {
4295        let a: Vec<f32> = (0..768).map(|i| (i as f32) * 0.01).collect();
4296        let b: Vec<f32> = (0..768).map(|i| (i as f32) * 0.02 + 0.5).collect();
4297        let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4298
4299        let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4300        let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4301        assert!(
4302            (dot - expected_dot).abs() < 1.0,
4303            "dot: expected {}, got {}",
4304            expected_dot,
4305            dot
4306        );
4307        assert!(
4308            (norm_b - expected_norm).abs() < 1.0,
4309            "norm: expected {}, got {}",
4310            expected_norm,
4311            norm_b
4312        );
4313    }
4314
4315    #[test]
4316    fn test_batch_cosine_scores() {
4317        // 4 vectors of dim 3
4318        let query = vec![1.0f32, 0.0, 0.0];
4319        let vectors = vec![
4320            1.0, 0.0, 0.0, // identical to query
4321            0.0, 1.0, 0.0, // orthogonal
4322            -1.0, 0.0, 0.0, // opposite
4323            0.5, 0.5, 0.0, // 45 degrees
4324        ];
4325        let mut scores = vec![0f32; 4];
4326        batch_cosine_scores(&query, &vectors, 3, &mut scores);
4327
4328        assert!((scores[0] - 1.0).abs() < 1e-5, "identical: {}", scores[0]);
4329        assert!(scores[1].abs() < 1e-5, "orthogonal: {}", scores[1]);
4330        assert!((scores[2] - (-1.0)).abs() < 1e-5, "opposite: {}", scores[2]);
4331        let expected_45 = 0.5f32 / (0.5f32.powi(2) + 0.5f32.powi(2)).sqrt();
4332        assert!(
4333            (scores[3] - expected_45).abs() < 1e-5,
4334            "45deg: expected {}, got {}",
4335            expected_45,
4336            scores[3]
4337        );
4338    }
4339
4340    #[test]
4341    fn test_batch_cosine_scores_matches_individual() {
4342        let query: Vec<f32> = (0..128).map(|i| (i as f32) * 0.1).collect();
4343        let n = 50;
4344        let dim = 128;
4345        let vectors: Vec<f32> = (0..n * dim).map(|i| ((i * 7 + 3) as f32) * 0.01).collect();
4346
4347        let mut batch_scores = vec![0f32; n];
4348        batch_cosine_scores(&query, &vectors, dim, &mut batch_scores);
4349
4350        for i in 0..n {
4351            let vec_i = &vectors[i * dim..(i + 1) * dim];
4352            let individual = cosine_similarity(&query, vec_i);
4353            assert!(
4354                (batch_scores[i] - individual).abs() < 1e-5,
4355                "vec {}: batch={}, individual={}",
4356                i,
4357                batch_scores[i],
4358                individual
4359            );
4360        }
4361    }
4362
4363    #[test]
4364    fn test_batch_cosine_scores_empty() {
4365        let query = vec![1.0f32, 2.0, 3.0];
4366        let vectors: Vec<f32> = vec![];
4367        let mut scores: Vec<f32> = vec![];
4368        batch_cosine_scores(&query, &vectors, 3, &mut scores);
4369        assert!(scores.is_empty());
4370    }
4371
4372    #[test]
4373    fn test_batch_cosine_scores_zero_query() {
4374        let query = vec![0.0f32, 0.0, 0.0];
4375        let vectors = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
4376        let mut scores = vec![0f32; 2];
4377        batch_cosine_scores(&query, &vectors, 3, &mut scores);
4378        assert_eq!(scores[0], 0.0);
4379        assert_eq!(scores[1], 0.0);
4380    }
4381
4382    // ================================================================
4383    // f16 conversion tests
4384    // ================================================================
4385
4386    #[test]
4387    fn test_f16_roundtrip_normal() {
4388        for &v in &[0.0f32, 1.0, -1.0, 0.5, -0.5, 0.333, 65504.0] {
4389            let h = f32_to_f16(v);
4390            let back = f16_to_f32(h);
4391            let err = (back - v).abs() / v.abs().max(1e-6);
4392            assert!(
4393                err < 0.002,
4394                "f16 roundtrip {v} → {h:#06x} → {back}, rel err {err}"
4395            );
4396        }
4397    }
4398
4399    #[test]
4400    fn test_f16_special() {
4401        // Zero
4402        assert_eq!(f16_to_f32(f32_to_f16(0.0)), 0.0);
4403        // Negative zero
4404        assert_eq!(f32_to_f16(-0.0), 0x8000);
4405        // Infinity
4406        assert!(f16_to_f32(f32_to_f16(f32::INFINITY)).is_infinite());
4407        // NaN
4408        assert!(f16_to_f32(f32_to_f16(f32::NAN)).is_nan());
4409    }
4410
4411    #[test]
4412    fn test_f16_embedding_range() {
4413        // Typical embedding values in [-1, 1]
4414        let values: Vec<f32> = (-100..=100).map(|i| i as f32 / 100.0).collect();
4415        for &v in &values {
4416            let back = f16_to_f32(f32_to_f16(v));
4417            assert!((back - v).abs() < 0.001, "f16 error for {v}: got {back}");
4418        }
4419    }
4420
4421    // ================================================================
4422    // u8 conversion tests
4423    // ================================================================
4424
4425    #[test]
4426    fn test_u8_roundtrip() {
4427        // Boundary values
4428        assert_eq!(f32_to_u8_saturating(-1.0), 0);
4429        assert_eq!(f32_to_u8_saturating(1.0), 255);
4430        assert_eq!(f32_to_u8_saturating(0.0), 127); // ~127.5 truncated
4431
4432        // Saturation
4433        assert_eq!(f32_to_u8_saturating(-2.0), 0);
4434        assert_eq!(f32_to_u8_saturating(2.0), 255);
4435    }
4436
4437    #[test]
4438    fn test_u8_dequantize() {
4439        assert!((u8_to_f32(0) - (-1.0)).abs() < 0.01);
4440        assert!((u8_to_f32(255) - 1.0).abs() < 0.01);
4441        assert!((u8_to_f32(127) - 0.0).abs() < 0.01);
4442    }
4443
4444    // ================================================================
4445    // Batch scoring tests for quantized vectors
4446    // ================================================================
4447
4448    #[test]
4449    fn test_batch_cosine_scores_f16() {
4450        let query = vec![0.6f32, 0.8, 0.0, 0.0];
4451        let dim = 4;
4452        let vecs_f32 = vec![
4453            0.6f32, 0.8, 0.0, 0.0, // identical to query
4454            0.0, 0.0, 0.6, 0.8, // orthogonal
4455        ];
4456
4457        // Quantize to f16
4458        let mut f16_buf = vec![0u16; 8];
4459        batch_f32_to_f16(&vecs_f32, &mut f16_buf);
4460        let raw: &[u8] =
4461            unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4462
4463        let mut scores = vec![0f32; 2];
4464        batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4465
4466        assert!(
4467            (scores[0] - 1.0).abs() < 0.01,
4468            "identical vectors: {}",
4469            scores[0]
4470        );
4471        assert!(scores[1].abs() < 0.01, "orthogonal vectors: {}", scores[1]);
4472    }
4473
4474    #[test]
4475    fn test_batch_cosine_scores_u8() {
4476        let query = vec![0.6f32, 0.8, 0.0, 0.0];
4477        let dim = 4;
4478        let vecs_f32 = vec![
4479            0.6f32, 0.8, 0.0, 0.0, // ~identical to query
4480            -0.6, -0.8, 0.0, 0.0, // opposite
4481        ];
4482
4483        // Quantize to u8
4484        let mut u8_buf = vec![0u8; 8];
4485        batch_f32_to_u8(&vecs_f32, &mut u8_buf);
4486
4487        let mut scores = vec![0f32; 2];
4488        batch_cosine_scores_u8(&query, &u8_buf, dim, &mut scores);
4489
4490        assert!(scores[0] > 0.95, "similar vectors: {}", scores[0]);
4491        assert!(scores[1] < -0.95, "opposite vectors: {}", scores[1]);
4492    }
4493
4494    #[test]
4495    fn test_batch_cosine_scores_f16_large_dim() {
4496        // Test with typical embedding dimension
4497        let dim = 768;
4498        let query: Vec<f32> = (0..dim).map(|i| (i as f32 / dim as f32) - 0.5).collect();
4499        let vec2: Vec<f32> = query.iter().map(|x| x * 0.9 + 0.01).collect();
4500
4501        let mut all_vecs = query.clone();
4502        all_vecs.extend_from_slice(&vec2);
4503
4504        let mut f16_buf = vec![0u16; all_vecs.len()];
4505        batch_f32_to_f16(&all_vecs, &mut f16_buf);
4506        let raw: &[u8] =
4507            unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4508
4509        let mut scores = vec![0f32; 2];
4510        batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4511
4512        // Self-similarity should be ~1.0
4513        assert!((scores[0] - 1.0).abs() < 0.01, "self-sim: {}", scores[0]);
4514        // High similarity with scaled version
4515        assert!(scores[1] > 0.99, "scaled-sim: {}", scores[1]);
4516    }
4517
4518    // ================================================================
4519    // Hamming distance tests
4520    // ================================================================
4521
4522    #[test]
4523    fn test_hamming_distance_identical() {
4524        let a = vec![0xAA; 64];
4525        assert_eq!(hamming_distance(&a, &a), 0);
4526    }
4527
4528    #[test]
4529    fn test_hamming_distance_opposite() {
4530        let a = vec![0xFF; 32];
4531        let b = vec![0x00; 32];
4532        assert_eq!(hamming_distance(&a, &b), 256);
4533    }
4534
4535    #[test]
4536    fn test_hamming_distance_known() {
4537        // Single byte: 0b10101010 vs 0b01010101 = 8 bits differ
4538        let a = vec![0xAA];
4539        let b = vec![0x55];
4540        assert_eq!(hamming_distance(&a, &b), 8);
4541
4542        // Two bytes
4543        let a = vec![0xFF, 0x00];
4544        let b = vec![0x00, 0x00];
4545        assert_eq!(hamming_distance(&a, &b), 8);
4546    }
4547
4548    #[test]
4549    fn test_hamming_distance_single_bit() {
4550        let a = vec![0x00; 16];
4551        let mut b = vec![0x00; 16];
4552        b[7] = 0x01; // flip one bit
4553        assert_eq!(hamming_distance(&a, &b), 1);
4554    }
4555
4556    #[test]
4557    fn test_hamming_distance_empty() {
4558        let a: Vec<u8> = vec![];
4559        assert_eq!(hamming_distance(&a, &a), 0);
4560    }
4561
4562    #[test]
4563    fn test_hamming_distance_remainder_path() {
4564        // 17 bytes: not aligned to 16 (NEON) or 32 (AVX2)
4565        let a = vec![0xFF; 17];
4566        let b = vec![0x00; 17];
4567        assert_eq!(hamming_distance(&a, &b), 136); // 17 * 8
4568
4569        // 33 bytes: tests 32-byte chunk + 1 remainder for AVX2
4570        let a = vec![0xFF; 33];
4571        let b = vec![0x00; 33];
4572        assert_eq!(hamming_distance(&a, &b), 264); // 33 * 8
4573    }
4574
4575    #[test]
4576    fn test_hamming_distance_large() {
4577        // 4096 bytes = 32768 bits, all differing
4578        let a = vec![0xFF; 4096];
4579        let b = vec![0x00; 4096];
4580        assert_eq!(hamming_distance(&a, &b), 32768);
4581    }
4582
4583    #[test]
4584    fn test_hamming_distance_scalar_matches() {
4585        // Verify SIMD path matches scalar for various sizes
4586        for size in [1, 7, 8, 15, 16, 31, 32, 63, 64, 100, 128, 255, 256] {
4587            let a: Vec<u8> = (0..size).map(|i| (i * 37 + 13) as u8).collect();
4588            let b: Vec<u8> = (0..size).map(|i| (i * 53 + 7) as u8).collect();
4589            let expected = hamming_distance_scalar(&a, &b);
4590            let got = hamming_distance(&a, &b);
4591            assert_eq!(got, expected, "mismatch at size {size}");
4592        }
4593    }
4594
4595    // ================================================================
4596    // Batch Hamming scoring tests
4597    // ================================================================
4598
4599    #[test]
4600    fn test_batch_hamming_scores_identical() {
4601        let query = vec![0xAA; 16];
4602        let db = vec![0xAA; 16]; // one vector, identical
4603        let mut scores = vec![0f32; 1];
4604        batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4605        assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4606    }
4607
4608    #[test]
4609    fn test_batch_hamming_scores_opposite() {
4610        let query = vec![0xFF; 16];
4611        let db = vec![0x00; 16];
4612        let mut scores = vec![0f32; 1];
4613        batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4614        assert!((scores[0] - 0.0).abs() < 1e-6, "opposite: {}", scores[0]);
4615    }
4616
4617    #[test]
4618    fn test_batch_hamming_scores_multiple() {
4619        let byte_len = 8;
4620        let dim_bits = 64;
4621        let query = vec![0xFF; byte_len];
4622        let mut db = Vec::new();
4623        db.extend_from_slice(&vec![0xFF; byte_len]); // identical → 1.0
4624        db.extend_from_slice(&vec![0x00; byte_len]); // opposite → 0.0
4625        db.extend_from_slice(&vec![0x0F; byte_len]); // half bits differ → 0.5
4626
4627        let mut scores = vec![0f32; 3];
4628        batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut scores);
4629
4630        assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4631        assert!((scores[1] - 0.0).abs() < 1e-6, "opposite: {}", scores[1]);
4632        assert!((scores[2] - 0.5).abs() < 1e-6, "half: {}", scores[2]);
4633    }
4634
4635    #[test]
4636    fn test_batch_hamming_scores_empty() {
4637        let query = vec![0xFF; 8];
4638        let db: Vec<u8> = vec![];
4639        let mut scores: Vec<f32> = vec![];
4640        batch_hamming_scores(&query, &db, 8, 64, &mut scores);
4641        assert!(scores.is_empty());
4642    }
4643
4644    #[test]
4645    fn test_batch_hamming_scores_zero_byte_len() {
4646        let query: Vec<u8> = vec![];
4647        let db: Vec<u8> = vec![];
4648        let mut scores = vec![0f32; 1];
4649        batch_hamming_scores(&query, &db, 0, 0, &mut scores);
4650        // Should return early without modifying scores
4651        assert_eq!(scores[0], 0.0);
4652    }
4653
4654    // ================================================================
4655    // Resolved-kernel and row-batched Hamming tests
4656    // ================================================================
4657
4658    fn hamming_matrix(rows: usize, byte_len: usize) -> (Vec<u8>, Vec<u8>) {
4659        let query: Vec<u8> = (0..byte_len).map(|i| (i * 31 + 5) as u8).collect();
4660        let db: Vec<u8> = (0..rows * byte_len)
4661            .map(|i| (i * 97 + i / byte_len * 11 + 3) as u8)
4662            .collect();
4663        (query, db)
4664    }
4665
4666    /// The row-batched kernels share query loads and accumulators across four
4667    /// rows; every width must still agree bit-for-bit with the scalar loop.
4668    #[test]
4669    fn batched_hamming_distances_match_scalar_for_every_row_count() {
4670        let kernel = HammingKernel::resolve();
4671        // Cover both multiples of the quad width and every tail remainder, and
4672        // byte lengths that exercise 16/32/64-byte chunking plus odd tails.
4673        for byte_len in [1, 7, 8, 15, 16, 31, 32, 33, 63, 64, 65, 128, 320] {
4674            for rows in [1, 2, 3, 4, 5, 7, 8, 9, 64, 70] {
4675                let (query, db) = hamming_matrix(rows, byte_len);
4676                let mut got = vec![0u32; rows];
4677                kernel.distances(&query, &db, byte_len, &mut got);
4678                for (row, &distance) in got.iter().enumerate() {
4679                    let expected =
4680                        hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4681                    assert_eq!(
4682                        distance, expected,
4683                        "row {row} of {rows} at byte_len {byte_len}"
4684                    );
4685                }
4686            }
4687        }
4688    }
4689
4690    #[test]
4691    fn gathered_hamming_distances_follow_row_ids() {
4692        let kernel = HammingKernel::resolve();
4693        let byte_len = 320;
4694        let rows = 37;
4695        let (query, db) = hamming_matrix(rows, byte_len);
4696        // Scattered, repeated and reversed ids: routing visits rows in graph
4697        // order, not storage order.
4698        let ids: Vec<u32> = [36, 0, 17, 17, 5, 31, 2, 9, 9, 36, 1].into_iter().collect();
4699        let mut got = vec![0u32; ids.len()];
4700        kernel.gather_distances(&query, &db, byte_len, &ids, &mut got);
4701        for (slot, &id) in ids.iter().enumerate() {
4702            let start = id as usize * byte_len;
4703            let expected = hamming_distance_scalar(&query, &db[start..start + byte_len]);
4704            assert_eq!(got[slot], expected, "slot {slot} for row {id}");
4705        }
4706    }
4707
4708    #[test]
4709    fn resolved_kernel_matches_scalar_pairwise() {
4710        let kernel = HammingKernel::resolve();
4711        for byte_len in [1, 8, 32, 64, 65, 320, 4096] {
4712            let (query, db) = hamming_matrix(1, byte_len);
4713            assert_eq!(
4714                kernel.distance(&query, &db),
4715                hamming_distance_scalar(&query, &db),
4716                "byte_len {byte_len}"
4717            );
4718        }
4719    }
4720
4721    #[test]
4722    fn scores_from_hamming_matches_batch_scores_across_blocks() {
4723        let kernel = HammingKernel::resolve();
4724        let byte_len = 320;
4725        let dim_bits = byte_len * 8;
4726        // More rows than one stack block so the block seam is covered.
4727        let rows = HAMMING_DISTANCE_BLOCK * 2 + 3;
4728        let (query, db) = hamming_matrix(rows, byte_len);
4729        let mut expected = vec![0f32; rows];
4730        for (row, score) in expected.iter_mut().enumerate() {
4731            let distance =
4732                hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4733            *score = 1.0 - distance as f32 / dim_bits as f32;
4734        }
4735        let mut got = vec![0f32; rows];
4736        scores_from_hamming(kernel, &query, &db, byte_len, dim_bits, &mut got);
4737        for (row, (&got, &want)) in got.iter().zip(expected.iter()).enumerate() {
4738            assert!((got - want).abs() < 1e-6, "row {row}: {got} vs {want}");
4739        }
4740        let mut public = vec![0f32; rows];
4741        batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut public);
4742        assert_eq!(got, public);
4743    }
4744}
4745
4746// ============================================================================
4747// SIMD-accelerated linear scan for sorted u32 slices (within-block seek)
4748// ============================================================================
4749
4750/// Find index of first element >= `target` in a sorted `u32` slice.
4751///
4752/// Equivalent to `slice.partition_point(|&d| d < target)` but uses SIMD to
4753/// scan 4 elements per cycle. Faster than binary search for slices ≤ 256
4754/// elements because it avoids the data-dependency chain inherent in binary
4755/// search (~8-10 cycles/iteration vs ~1-2 cycles/iteration for SIMD scan).
4756///
4757/// Returns `slice.len()` if no element >= `target`.
4758#[inline]
4759pub fn find_first_ge_u32(slice: &[u32], target: u32) -> usize {
4760    #[cfg(target_arch = "aarch64")]
4761    {
4762        if neon::is_available() {
4763            return unsafe { find_first_ge_u32_neon(slice, target) };
4764        }
4765    }
4766
4767    #[cfg(target_arch = "x86_64")]
4768    {
4769        if sse::is_available() {
4770            return unsafe { find_first_ge_u32_sse(slice, target) };
4771        }
4772    }
4773
4774    // Scalar fallback (WASM, other architectures)
4775    slice.partition_point(|&d| d < target)
4776}
4777
4778#[cfg(target_arch = "aarch64")]
4779#[target_feature(enable = "neon")]
4780#[allow(unsafe_op_in_unsafe_fn)]
4781unsafe fn find_first_ge_u32_neon(slice: &[u32], target: u32) -> usize {
4782    use std::arch::aarch64::*;
4783
4784    let n = slice.len();
4785    let ptr = slice.as_ptr();
4786    let target_vec = vdupq_n_u32(target);
4787    // Bit positions for each lane: [1, 2, 4, 8]
4788    let bit_mask: uint32x4_t = core::mem::transmute([1u32, 2u32, 4u32, 8u32]);
4789
4790    let chunks = n / 16;
4791    let mut base = 0usize;
4792
4793    // Process 16 elements per iteration (4 × 4-wide NEON compares)
4794    for _ in 0..chunks {
4795        let v0 = vld1q_u32(ptr.add(base));
4796        let v1 = vld1q_u32(ptr.add(base + 4));
4797        let v2 = vld1q_u32(ptr.add(base + 8));
4798        let v3 = vld1q_u32(ptr.add(base + 12));
4799
4800        let c0 = vcgeq_u32(v0, target_vec);
4801        let c1 = vcgeq_u32(v1, target_vec);
4802        let c2 = vcgeq_u32(v2, target_vec);
4803        let c3 = vcgeq_u32(v3, target_vec);
4804
4805        let m0 = vaddvq_u32(vandq_u32(c0, bit_mask));
4806        if m0 != 0 {
4807            return base + m0.trailing_zeros() as usize;
4808        }
4809        let m1 = vaddvq_u32(vandq_u32(c1, bit_mask));
4810        if m1 != 0 {
4811            return base + 4 + m1.trailing_zeros() as usize;
4812        }
4813        let m2 = vaddvq_u32(vandq_u32(c2, bit_mask));
4814        if m2 != 0 {
4815            return base + 8 + m2.trailing_zeros() as usize;
4816        }
4817        let m3 = vaddvq_u32(vandq_u32(c3, bit_mask));
4818        if m3 != 0 {
4819            return base + 12 + m3.trailing_zeros() as usize;
4820        }
4821        base += 16;
4822    }
4823
4824    // Process remaining 4 elements at a time
4825    while base + 4 <= n {
4826        let vals = vld1q_u32(ptr.add(base));
4827        let cmp = vcgeq_u32(vals, target_vec);
4828        let mask = vaddvq_u32(vandq_u32(cmp, bit_mask));
4829        if mask != 0 {
4830            return base + mask.trailing_zeros() as usize;
4831        }
4832        base += 4;
4833    }
4834
4835    // Scalar remainder (0-3 elements)
4836    while base < n {
4837        if *slice.get_unchecked(base) >= target {
4838            return base;
4839        }
4840        base += 1;
4841    }
4842    n
4843}
4844
4845#[cfg(target_arch = "x86_64")]
4846#[target_feature(enable = "sse2")]
4847#[allow(unsafe_op_in_unsafe_fn)]
4848unsafe fn find_first_ge_u32_sse(slice: &[u32], target: u32) -> usize {
4849    use std::arch::x86_64::*;
4850
4851    let n = slice.len();
4852    let ptr = slice.as_ptr();
4853
4854    // For unsigned >= comparison: XOR with 0x80000000 converts to signed domain
4855    let sign_flip = _mm_set1_epi32(i32::MIN);
4856    let target_xor = _mm_xor_si128(_mm_set1_epi32(target as i32), sign_flip);
4857
4858    let chunks = n / 16;
4859    let mut base = 0usize;
4860
4861    // Process 16 elements per iteration (4 × 4-wide SSE compares)
4862    for _ in 0..chunks {
4863        let v0 = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4864        let v1 = _mm_xor_si128(
4865            _mm_loadu_si128(ptr.add(base + 4) as *const __m128i),
4866            sign_flip,
4867        );
4868        let v2 = _mm_xor_si128(
4869            _mm_loadu_si128(ptr.add(base + 8) as *const __m128i),
4870            sign_flip,
4871        );
4872        let v3 = _mm_xor_si128(
4873            _mm_loadu_si128(ptr.add(base + 12) as *const __m128i),
4874            sign_flip,
4875        );
4876
4877        // ge = eq | gt (in signed domain after XOR)
4878        let ge0 = _mm_or_si128(
4879            _mm_cmpeq_epi32(v0, target_xor),
4880            _mm_cmpgt_epi32(v0, target_xor),
4881        );
4882        let m0 = _mm_movemask_ps(_mm_castsi128_ps(ge0)) as u32;
4883        if m0 != 0 {
4884            return base + m0.trailing_zeros() as usize;
4885        }
4886
4887        let ge1 = _mm_or_si128(
4888            _mm_cmpeq_epi32(v1, target_xor),
4889            _mm_cmpgt_epi32(v1, target_xor),
4890        );
4891        let m1 = _mm_movemask_ps(_mm_castsi128_ps(ge1)) as u32;
4892        if m1 != 0 {
4893            return base + 4 + m1.trailing_zeros() as usize;
4894        }
4895
4896        let ge2 = _mm_or_si128(
4897            _mm_cmpeq_epi32(v2, target_xor),
4898            _mm_cmpgt_epi32(v2, target_xor),
4899        );
4900        let m2 = _mm_movemask_ps(_mm_castsi128_ps(ge2)) as u32;
4901        if m2 != 0 {
4902            return base + 8 + m2.trailing_zeros() as usize;
4903        }
4904
4905        let ge3 = _mm_or_si128(
4906            _mm_cmpeq_epi32(v3, target_xor),
4907            _mm_cmpgt_epi32(v3, target_xor),
4908        );
4909        let m3 = _mm_movemask_ps(_mm_castsi128_ps(ge3)) as u32;
4910        if m3 != 0 {
4911            return base + 12 + m3.trailing_zeros() as usize;
4912        }
4913        base += 16;
4914    }
4915
4916    // Process remaining 4 elements at a time
4917    while base + 4 <= n {
4918        let vals = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4919        let ge = _mm_or_si128(
4920            _mm_cmpeq_epi32(vals, target_xor),
4921            _mm_cmpgt_epi32(vals, target_xor),
4922        );
4923        let mask = _mm_movemask_ps(_mm_castsi128_ps(ge)) as u32;
4924        if mask != 0 {
4925            return base + mask.trailing_zeros() as usize;
4926        }
4927        base += 4;
4928    }
4929
4930    // Scalar remainder (0-3 elements)
4931    while base < n {
4932        if *slice.get_unchecked(base) >= target {
4933            return base;
4934        }
4935        base += 1;
4936    }
4937    n
4938}
4939
4940#[cfg(test)]
4941mod find_first_ge_tests {
4942    use super::find_first_ge_u32;
4943
4944    #[test]
4945    fn test_find_first_ge_basic() {
4946        let data: Vec<u32> = (0..128).map(|i| i * 3).collect(); // [0, 3, 6, ..., 381]
4947        assert_eq!(find_first_ge_u32(&data, 0), 0);
4948        assert_eq!(find_first_ge_u32(&data, 1), 1); // first >= 1 is 3 at idx 1
4949        assert_eq!(find_first_ge_u32(&data, 3), 1);
4950        assert_eq!(find_first_ge_u32(&data, 4), 2); // first >= 4 is 6 at idx 2
4951        assert_eq!(find_first_ge_u32(&data, 381), 127);
4952        assert_eq!(find_first_ge_u32(&data, 382), 128); // past end
4953    }
4954
4955    #[test]
4956    fn test_find_first_ge_matches_partition_point() {
4957        let data: Vec<u32> = vec![1, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55, 60, 65, 70, 75];
4958        for target in 0..80 {
4959            let expected = data.partition_point(|&d| d < target);
4960            let actual = find_first_ge_u32(&data, target);
4961            assert_eq!(actual, expected, "target={}", target);
4962        }
4963    }
4964
4965    #[test]
4966    fn test_find_first_ge_small_slices() {
4967        // Empty
4968        assert_eq!(find_first_ge_u32(&[], 5), 0);
4969        // Single element
4970        assert_eq!(find_first_ge_u32(&[10], 5), 0);
4971        assert_eq!(find_first_ge_u32(&[10], 10), 0);
4972        assert_eq!(find_first_ge_u32(&[10], 11), 1);
4973        // Three elements (< SIMD width)
4974        assert_eq!(find_first_ge_u32(&[2, 4, 6], 5), 2);
4975    }
4976
4977    #[test]
4978    fn test_find_first_ge_full_block() {
4979        // Simulate a full 128-entry block
4980        let data: Vec<u32> = (100..228).collect();
4981        assert_eq!(find_first_ge_u32(&data, 100), 0);
4982        assert_eq!(find_first_ge_u32(&data, 150), 50);
4983        assert_eq!(find_first_ge_u32(&data, 227), 127);
4984        assert_eq!(find_first_ge_u32(&data, 228), 128);
4985        assert_eq!(find_first_ge_u32(&data, 99), 0);
4986    }
4987
4988    #[test]
4989    fn test_find_first_ge_u32_max() {
4990        // Test with large u32 values (unsigned correctness)
4991        let data = vec![u32::MAX - 10, u32::MAX - 5, u32::MAX - 1, u32::MAX];
4992        assert_eq!(find_first_ge_u32(&data, u32::MAX - 10), 0);
4993        assert_eq!(find_first_ge_u32(&data, u32::MAX - 7), 1);
4994        assert_eq!(find_first_ge_u32(&data, u32::MAX), 3);
4995    }
4996}
4997
4998/// Regression coverage for the algebraic (reassociation-permitting) float
4999/// reductions. Pins them against a strict f64 reference and against the
5000/// hand-written SIMD kernels they must stay interchangeable with.
5001#[cfg(test)]
5002mod algebraic_reduction_tests {
5003    use super::*;
5004
5005    fn algebraic_test_vector(dim: usize, seed: u64) -> Vec<f32> {
5006        let mut state = seed | 1;
5007        (0..dim)
5008            .map(|_| {
5009                state = state
5010                    .wrapping_mul(6364136223846793005)
5011                    .wrapping_add(1442695040888963407);
5012                ((state >> 33) as f32 / (1u64 << 30) as f32) - 1.0
5013            })
5014            .collect()
5015    }
5016
5017    /// Dimensions that straddle every SIMD chunk width used in this module
5018    /// (4/8/16 lanes) plus their tails, and realistic embedding widths.
5019    const ALGEBRAIC_TEST_DIMS: [usize; 12] = [0, 1, 3, 4, 7, 8, 15, 16, 17, 64, 384, 768];
5020
5021    #[test]
5022    fn test_algebraic_squared_l2_matches_f64_reference() {
5023        for dim in ALGEBRAIC_TEST_DIMS {
5024            let a = algebraic_test_vector(dim, 0x51ed_0001);
5025            let b = algebraic_test_vector(dim, 0x51ed_0002);
5026            let reference: f64 = a
5027                .iter()
5028                .zip(&b)
5029                .map(|(&x, &y)| {
5030                    let delta = f64::from(x) - f64::from(y);
5031                    delta * delta
5032                })
5033                .sum();
5034            let actual = squared_l2_f32(&a, &b);
5035            let tolerance = (reference * 1e-5).max(1e-6);
5036            assert!(
5037                (f64::from(actual) - reference).abs() <= tolerance,
5038                "dim {dim}: squared_l2_f32 {actual} drifted from f64 reference {reference}"
5039            );
5040        }
5041    }
5042
5043    #[test]
5044    fn test_algebraic_squared_l2_uses_shorter_length() {
5045        let a = [1.0f32, 2.0, 3.0, 4.0];
5046        let b = [1.0f32, 4.0];
5047        assert_eq!(squared_l2_f32(&a, &b), 4.0);
5048        assert_eq!(squared_l2_f32(&b, &a), 4.0);
5049    }
5050
5051    #[test]
5052    fn test_algebraic_norm_squared_matches_f64_reference() {
5053        for dim in ALGEBRAIC_TEST_DIMS {
5054            let v = algebraic_test_vector(dim, 0x51ed_0003);
5055            let reference: f64 = v.iter().map(|&x| f64::from(x) * f64::from(x)).sum();
5056            let actual = norm_squared_f32(&v);
5057            let tolerance = (reference * 1e-5).max(1e-6);
5058            assert!(
5059                (f64::from(actual) - reference).abs() <= tolerance,
5060                "dim {dim}: norm_squared_f32 {actual} drifted from f64 reference {reference}"
5061            );
5062            assert!(
5063                (f64::from(norm_f32(&v)) - reference.sqrt()).abs() <= tolerance.sqrt().max(1e-5)
5064            );
5065        }
5066    }
5067
5068    #[test]
5069    fn test_algebraic_dot_scalar_matches_simd_dispatch() {
5070        for dim in ALGEBRAIC_TEST_DIMS {
5071            let a = algebraic_test_vector(dim, 0x51ed_0004);
5072            let b = algebraic_test_vector(dim, 0x51ed_0005);
5073            let dispatched = dot_product_f32(&a, &b, dim);
5074            let scalar = dot_product_f32_scalar(&a, &b);
5075            let tolerance = (dispatched.abs() * 1e-5).max(1e-5);
5076            assert!(
5077                (dispatched - scalar).abs() <= tolerance,
5078                "dim {dim}: scalar dot {scalar} disagrees with dispatched {dispatched}"
5079            );
5080
5081            let (fused_dot, fused_norm) = fused_dot_norm(&a, &b, dim);
5082            let (scalar_dot, scalar_norm) = fused_dot_norm_scalar(&a, &b);
5083            assert!((fused_dot - scalar_dot).abs() <= tolerance);
5084            assert!((fused_norm - scalar_norm).abs() <= (fused_norm.abs() * 1e-5).max(1e-5));
5085        }
5086    }
5087
5088    /// Rust's algebraic operations differ from `-ffast-math`: they permit
5089    /// reassociation but never assume finite inputs. NaN and infinity must
5090    /// still propagate, otherwise a degenerate stored vector would silently
5091    /// score as a finite number instead of being rejected downstream.
5092    #[test]
5093    fn test_algebraic_reductions_propagate_non_finite() {
5094        let finite = vec![1.0f32; 8];
5095
5096        let mut with_nan = finite.clone();
5097        with_nan[5] = f32::NAN;
5098        assert!(norm_squared_f32(&with_nan).is_nan());
5099        assert!(squared_l2_f32(&with_nan, &finite).is_nan());
5100        assert!(dot_product_f32_scalar(&with_nan, &finite).is_nan());
5101
5102        let mut with_inf = finite.clone();
5103        with_inf[2] = f32::INFINITY;
5104        assert!(norm_squared_f32(&with_inf).is_infinite());
5105        assert!(squared_l2_f32(&with_inf, &finite).is_infinite());
5106        assert!(dot_product_f32_scalar(&with_inf, &finite).is_infinite());
5107    }
5108}