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