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 f32 slice to u8 with [-1,1] → [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 [-1,1]→[0,255]).
3111/// Converts u8→f32 using NEON widening chain (16 values/iteration), scores with FMA.
3112/// Memory bandwidth is quartered compared to f32 scoring.
3113#[inline]
3114pub fn batch_cosine_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3115    let n = scores.len();
3116    let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3117    assert_eq!(query.len(), dim, "u8 batch cosine query dimension mismatch");
3118    assert!(
3119        vectors_raw.len() >= required,
3120        "u8 batch cosine vectors are truncated: need {required} bytes, got {}",
3121        vectors_raw.len()
3122    );
3123    if dim == 0 || n == 0 {
3124        return;
3125    }
3126
3127    let norm_q_sq = dot_product_f32(query, query, dim);
3128    if norm_q_sq < f32::EPSILON {
3129        for s in scores.iter_mut() {
3130            *s = 0.0;
3131        }
3132        return;
3133    }
3134    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3135
3136    for i in 0..n {
3137        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3138
3139        let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3140        scores[i] = if norm_v_sq < f32::EPSILON {
3141            0.0
3142        } else {
3143            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3144        };
3145    }
3146}
3147
3148// ============================================================================
3149// Batch dot-product scoring for unit-norm vectors
3150// ============================================================================
3151
3152/// Batch dot-product scoring: f32 query vs N contiguous f32 unit-norm vectors.
3153///
3154/// For pre-normalized vectors (||v|| = 1), cosine = dot(q, v) / ||q||.
3155/// Skips per-vector norm computation — ~40% less work than `batch_cosine_scores`.
3156#[inline]
3157pub fn batch_dot_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
3158    let n = scores.len();
3159    let required = n
3160        .checked_mul(dim)
3161        .expect("batch dot vector length overflow");
3162    assert_eq!(query.len(), dim, "batch dot query dimension mismatch");
3163    assert!(
3164        vectors.len() >= required,
3165        "batch dot vectors are truncated: need {required}, got {}",
3166        vectors.len()
3167    );
3168
3169    if dim == 0 || n == 0 {
3170        return;
3171    }
3172
3173    let norm_q_sq = dot_product_f32(query, query, dim);
3174    if norm_q_sq < f32::EPSILON {
3175        for s in scores.iter_mut() {
3176            *s = 0.0;
3177        }
3178        return;
3179    }
3180    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3181
3182    for i in 0..n {
3183        let vec = &vectors[i * dim..(i + 1) * dim];
3184        let dot = dot_product_f32(query, vec, dim);
3185        scores[i] = dot * inv_norm_q;
3186    }
3187}
3188
3189/// Batch dot-product scoring: f32 query vs N contiguous f16 unit-norm vectors.
3190///
3191/// For pre-normalized vectors (||v|| = 1), cosine = dot(q, v) / ||q||.
3192/// Uses F16C/NEON hardware conversion + dot-only kernel.
3193#[inline]
3194pub fn batch_dot_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3195    let n = scores.len();
3196    let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3197    let required = n
3198        .checked_mul(vec_bytes)
3199        .expect("f16 batch byte length overflow");
3200    assert_eq!(query.len(), dim, "f16 batch dot query dimension mismatch");
3201    assert!(
3202        vectors_raw.len() >= required,
3203        "f16 batch dot vectors are truncated: need {required} bytes, got {}",
3204        vectors_raw.len()
3205    );
3206    if required > 0 {
3207        assert!(
3208            (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3209            "f16 batch dot vectors are not 2-byte aligned"
3210        );
3211    }
3212    if dim == 0 || n == 0 {
3213        return;
3214    }
3215
3216    let norm_q_sq = dot_product_f32(query, query, dim);
3217    if norm_q_sq < f32::EPSILON {
3218        for s in scores.iter_mut() {
3219            *s = 0.0;
3220        }
3221        return;
3222    }
3223    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3224
3225    let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3226    for i in 0..n {
3227        let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3228        let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3229        let dot = dot_product_f16_quant(&query_f16, f16_slice, dim);
3230        scores[i] = dot * inv_norm_q;
3231    }
3232}
3233
3234/// Batch dot-product scoring: f32 query vs N contiguous u8 unit-norm vectors.
3235///
3236/// For pre-normalized vectors (||v|| = 1), cosine = dot(q, v) / ||q||.
3237/// Uses NEON/SSE widening chain for u8→f32 conversion + dot-only kernel.
3238#[inline]
3239pub fn batch_dot_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3240    let n = scores.len();
3241    let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3242    assert_eq!(query.len(), dim, "u8 batch dot query dimension mismatch");
3243    assert!(
3244        vectors_raw.len() >= required,
3245        "u8 batch dot vectors are truncated: need {required} bytes, got {}",
3246        vectors_raw.len()
3247    );
3248    if dim == 0 || n == 0 {
3249        return;
3250    }
3251
3252    let norm_q_sq = dot_product_f32(query, query, dim);
3253    if norm_q_sq < f32::EPSILON {
3254        for s in scores.iter_mut() {
3255            *s = 0.0;
3256        }
3257        return;
3258    }
3259    let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3260
3261    for i in 0..n {
3262        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3263        let dot = dot_product_u8_quant(query, u8_slice, dim);
3264        scores[i] = dot * inv_norm_q;
3265    }
3266}
3267
3268// ============================================================================
3269// Precomputed-norm batch scoring (avoids redundant query norm + f16 conversion)
3270// ============================================================================
3271
3272/// Batch cosine: f32 query vs N f32 vectors, with precomputed `inv_norm_q`.
3273#[inline]
3274pub fn batch_cosine_scores_precomp(
3275    query: &[f32],
3276    vectors: &[f32],
3277    dim: usize,
3278    scores: &mut [f32],
3279    inv_norm_q: f32,
3280) {
3281    let n = scores.len();
3282    let required = n
3283        .checked_mul(dim)
3284        .expect("precomputed cosine vector length overflow");
3285    assert_eq!(
3286        query.len(),
3287        dim,
3288        "precomputed cosine query dimension mismatch"
3289    );
3290    assert!(
3291        vectors.len() >= required,
3292        "precomputed cosine vectors are truncated: need {required}, got {}",
3293        vectors.len()
3294    );
3295    for i in 0..n {
3296        let vec = &vectors[i * dim..(i + 1) * dim];
3297        let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
3298        scores[i] = if norm_v_sq < f32::EPSILON {
3299            0.0
3300        } else {
3301            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3302        };
3303    }
3304}
3305
3306/// Batch cosine: precomputed `inv_norm_q` + `query_f16` vs N f16 vectors.
3307#[inline]
3308pub fn batch_cosine_scores_f16_precomp(
3309    query_f16: &[u16],
3310    vectors_raw: &[u8],
3311    dim: usize,
3312    scores: &mut [f32],
3313    inv_norm_q: f32,
3314) {
3315    let n = scores.len();
3316    let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3317    let required = n
3318        .checked_mul(vec_bytes)
3319        .expect("precomputed f16 cosine batch byte length overflow");
3320    assert_eq!(
3321        query_f16.len(),
3322        dim,
3323        "precomputed f16 cosine query dimension mismatch"
3324    );
3325    assert!(
3326        vectors_raw.len() >= required,
3327        "precomputed f16 cosine vectors are truncated: need {required} bytes, got {}",
3328        vectors_raw.len()
3329    );
3330    if required > 0 {
3331        assert!(
3332            (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3333            "precomputed f16 cosine vectors are not 2-byte aligned"
3334        );
3335    }
3336    for i in 0..n {
3337        let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3338        let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3339        let (dot, norm_v_sq) = fused_dot_norm_f16(query_f16, f16_slice, dim);
3340        scores[i] = if norm_v_sq < f32::EPSILON {
3341            0.0
3342        } else {
3343            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3344        };
3345    }
3346}
3347
3348/// Batch cosine: precomputed `inv_norm_q` vs N u8 vectors.
3349#[inline]
3350pub fn batch_cosine_scores_u8_precomp(
3351    query: &[f32],
3352    vectors_raw: &[u8],
3353    dim: usize,
3354    scores: &mut [f32],
3355    inv_norm_q: f32,
3356) {
3357    let n = scores.len();
3358    let required = n
3359        .checked_mul(dim)
3360        .expect("precomputed u8 cosine batch byte length overflow");
3361    assert_eq!(
3362        query.len(),
3363        dim,
3364        "precomputed u8 cosine query dimension mismatch"
3365    );
3366    assert!(
3367        vectors_raw.len() >= required,
3368        "precomputed u8 cosine vectors are truncated: need {required} bytes, got {}",
3369        vectors_raw.len()
3370    );
3371    for i in 0..n {
3372        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3373        let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3374        scores[i] = if norm_v_sq < f32::EPSILON {
3375            0.0
3376        } else {
3377            dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3378        };
3379    }
3380}
3381
3382/// Batch dot-product: precomputed `inv_norm_q` vs N f32 unit-norm vectors.
3383#[inline]
3384pub fn batch_dot_scores_precomp(
3385    query: &[f32],
3386    vectors: &[f32],
3387    dim: usize,
3388    scores: &mut [f32],
3389    inv_norm_q: f32,
3390) {
3391    let n = scores.len();
3392    let required = n
3393        .checked_mul(dim)
3394        .expect("precomputed dot vector length overflow");
3395    assert_eq!(query.len(), dim, "precomputed dot query dimension mismatch");
3396    assert!(
3397        vectors.len() >= required,
3398        "precomputed dot vectors are truncated: need {required}, got {}",
3399        vectors.len()
3400    );
3401    for i in 0..n {
3402        let vec = &vectors[i * dim..(i + 1) * dim];
3403        scores[i] = dot_product_f32(query, vec, dim) * inv_norm_q;
3404    }
3405}
3406
3407/// Batch dot-product: precomputed `inv_norm_q` + `query_f16` vs N f16 unit-norm vectors.
3408#[inline]
3409pub fn batch_dot_scores_f16_precomp(
3410    query_f16: &[u16],
3411    vectors_raw: &[u8],
3412    dim: usize,
3413    scores: &mut [f32],
3414    inv_norm_q: f32,
3415) {
3416    let n = scores.len();
3417    let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3418    let required = n
3419        .checked_mul(vec_bytes)
3420        .expect("precomputed f16 dot batch byte length overflow");
3421    assert_eq!(
3422        query_f16.len(),
3423        dim,
3424        "precomputed f16 dot query dimension mismatch"
3425    );
3426    assert!(
3427        vectors_raw.len() >= required,
3428        "precomputed f16 dot vectors are truncated: need {required} bytes, got {}",
3429        vectors_raw.len()
3430    );
3431    if required > 0 {
3432        assert!(
3433            (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3434            "precomputed f16 dot vectors are not 2-byte aligned"
3435        );
3436    }
3437    for i in 0..n {
3438        let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3439        let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3440        scores[i] = dot_product_f16_quant(query_f16, f16_slice, dim) * inv_norm_q;
3441    }
3442}
3443
3444/// Batch dot-product: precomputed `inv_norm_q` vs N u8 unit-norm vectors.
3445#[inline]
3446pub fn batch_dot_scores_u8_precomp(
3447    query: &[f32],
3448    vectors_raw: &[u8],
3449    dim: usize,
3450    scores: &mut [f32],
3451    inv_norm_q: f32,
3452) {
3453    let n = scores.len();
3454    let required = n
3455        .checked_mul(dim)
3456        .expect("precomputed u8 dot batch byte length overflow");
3457    assert_eq!(
3458        query.len(),
3459        dim,
3460        "precomputed u8 dot query dimension mismatch"
3461    );
3462    assert!(
3463        vectors_raw.len() >= required,
3464        "precomputed u8 dot vectors are truncated: need {required} bytes, got {}",
3465        vectors_raw.len()
3466    );
3467    for i in 0..n {
3468        let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3469        scores[i] = dot_product_u8_quant(query, u8_slice, dim) * inv_norm_q;
3470    }
3471}
3472
3473/// Compute cosine similarity between two f32 vectors with SIMD acceleration
3474///
3475/// Returns dot(a,b) / (||a|| * ||b||), range [-1, 1]
3476/// Returns 0.0 if either vector has zero norm.
3477#[inline]
3478pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
3479    assert_eq!(a.len(), b.len(), "cosine vector dimension mismatch");
3480    let count = a.len();
3481
3482    if count == 0 {
3483        return 0.0;
3484    }
3485
3486    let dot = dot_product_f32(a, b, count);
3487    let norm_a = dot_product_f32(a, a, count);
3488    let norm_b = dot_product_f32(b, b, count);
3489
3490    let denom = (norm_a * norm_b).sqrt();
3491    if denom < f32::EPSILON {
3492        return 0.0;
3493    }
3494
3495    dot / denom
3496}
3497
3498// ============================================================================
3499// Hamming distance for binary dense vectors
3500// ============================================================================
3501
3502/// AVX-512 Hamming distance using `VPOPCNTDQ`.
3503///
3504/// Processes 64 bytes per iteration with a single hardware popcount per lane
3505/// group, which removes the nibble-lookup shuffles the AVX2 path needs.
3506#[cfg(target_arch = "x86_64")]
3507#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3508#[allow(unsafe_op_in_unsafe_fn)]
3509unsafe fn hamming_distance_avx512(a: &[u8], b: &[u8]) -> u32 {
3510    use std::arch::x86_64::*;
3511
3512    let len = a.len();
3513    let chunks64 = len / 64;
3514    let mut acc = _mm512_setzero_si512();
3515
3516    for c in 0..chunks64 {
3517        let off = c * 64;
3518        let va = _mm512_loadu_si512(a.as_ptr().add(off) as *const __m512i);
3519        let vb = _mm512_loadu_si512(b.as_ptr().add(off) as *const __m512i);
3520        acc = _mm512_add_epi64(acc, _mm512_popcnt_epi64(_mm512_xor_si512(va, vb)));
3521    }
3522
3523    let base = chunks64 * 64;
3524    _mm512_reduce_add_epi64(acc) as u32 + hamming_distance_scalar(&a[base..], &b[base..])
3525}
3526
3527/// Four-row AVX-512 Hamming distance sharing the query load across rows.
3528#[cfg(target_arch = "x86_64")]
3529#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3530#[allow(unsafe_op_in_unsafe_fn)]
3531unsafe fn hamming_distance_x4_avx512(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3532    use std::arch::x86_64::*;
3533
3534    let len = query.len();
3535    let chunks64 = len / 64;
3536    let mut acc = [_mm512_setzero_si512(); 4];
3537
3538    for c in 0..chunks64 {
3539        let off = c * 64;
3540        let vq = _mm512_loadu_si512(query.as_ptr().add(off) as *const __m512i);
3541        for r in 0..4 {
3542            let vr = _mm512_loadu_si512(rows[r].as_ptr().add(off) as *const __m512i);
3543            acc[r] = _mm512_add_epi64(acc[r], _mm512_popcnt_epi64(_mm512_xor_si512(vq, vr)));
3544        }
3545    }
3546
3547    let base = chunks64 * 64;
3548    let tail = &query[base..];
3549    [
3550        _mm512_reduce_add_epi64(acc[0]) as u32 + hamming_distance_scalar(tail, &rows[0][base..]),
3551        _mm512_reduce_add_epi64(acc[1]) as u32 + hamming_distance_scalar(tail, &rows[1][base..]),
3552        _mm512_reduce_add_epi64(acc[2]) as u32 + hamming_distance_scalar(tail, &rows[2][base..]),
3553        _mm512_reduce_add_epi64(acc[3]) as u32 + hamming_distance_scalar(tail, &rows[3][base..]),
3554    ]
3555}
3556
3557/// Four-row scalar Hamming distance sharing the query load across rows.
3558#[inline]
3559fn hamming_distance_x4_scalar(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3560    let len = query.len();
3561    let chunks = len / 8;
3562    let mut total = [0u32; 4];
3563
3564    for i in 0..chunks {
3565        let off = i * 8;
3566        let vq = unsafe { std::ptr::read_unaligned(query.as_ptr().add(off) as *const u64) };
3567        for r in 0..4 {
3568            let vr = unsafe { std::ptr::read_unaligned(rows[r].as_ptr().add(off) as *const u64) };
3569            total[r] += (vq ^ vr).count_ones();
3570        }
3571    }
3572
3573    let base = chunks * 8;
3574    for k in base..len {
3575        let q = query[k];
3576        for r in 0..4 {
3577            total[r] += (q ^ rows[r][k]).count_ones();
3578        }
3579    }
3580
3581    total
3582}
3583
3584/// Rows scored per kernel invocation. Sharing the query load, the AVX2 nibble
3585/// lookup table and the horizontal reduction across four rows amortises the
3586/// non-inlinable `#[target_feature]` call and overlaps the popcount chains.
3587const HAMMING_ROWS_PER_KERNEL: usize = 4;
3588
3589/// Architecture kernel resolved once for a whole scan.
3590///
3591/// Hot binary paths — HNSW centroid routing, k-majority assignment, leaf
3592/// scanning — score millions of code pairs against one query. Resolving the
3593/// kernel up front keeps runtime feature detection out of the inner loop, and
3594/// the row-batched entry points let one dispatch cover a whole neighbour list.
3595#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3596pub enum HammingKernel {
3597    #[cfg(target_arch = "x86_64")]
3598    Avx512,
3599    #[cfg(target_arch = "x86_64")]
3600    Avx2,
3601    #[cfg(target_arch = "aarch64")]
3602    Neon,
3603    Scalar,
3604}
3605
3606impl HammingKernel {
3607    /// Detect the widest kernel this CPU supports.
3608    #[inline]
3609    pub fn resolve() -> Self {
3610        #[cfg(target_arch = "x86_64")]
3611        {
3612            if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512vpopcntdq") {
3613                return Self::Avx512;
3614            }
3615            if avx2::is_available() {
3616                return Self::Avx2;
3617            }
3618            Self::Scalar
3619        }
3620
3621        #[cfg(target_arch = "aarch64")]
3622        {
3623            Self::Neon
3624        }
3625
3626        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
3627        {
3628            Self::Scalar
3629        }
3630    }
3631
3632    /// Hamming distance between two equal-length packed-bit vectors.
3633    #[inline]
3634    pub fn distance(self, a: &[u8], b: &[u8]) -> u32 {
3635        debug_assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3636        match self {
3637            #[cfg(target_arch = "x86_64")]
3638            Self::Avx512 => unsafe { hamming_distance_avx512(a, b) },
3639            #[cfg(target_arch = "x86_64")]
3640            Self::Avx2 => unsafe { avx2::hamming_distance(a, b) },
3641            #[cfg(target_arch = "aarch64")]
3642            Self::Neon => unsafe { neon::hamming_distance(a, b) },
3643            Self::Scalar => hamming_distance_scalar(a, b),
3644        }
3645    }
3646
3647    /// `out[i]` receives the distance from `query` to row `i` of `db`.
3648    pub fn distances(self, query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3649        self.score_rows(query, db, byte_len, out, |index| index);
3650    }
3651
3652    /// `out[i]` receives the distance from `query` to row `ids[i]` of `db`.
3653    ///
3654    /// Graph routing visits scattered centroid rows; gathering them through one
3655    /// dispatch keeps the batched kernel usable there.
3656    pub fn gather_distances(
3657        self,
3658        query: &[u8],
3659        db: &[u8],
3660        byte_len: usize,
3661        ids: &[u32],
3662        out: &mut [u32],
3663    ) {
3664        assert_eq!(
3665            ids.len(),
3666            out.len(),
3667            "Hamming gather needs one output slot per row id"
3668        );
3669        self.score_rows(query, db, byte_len, out, |index| ids[index] as usize);
3670    }
3671
3672    #[inline]
3673    fn score_rows(
3674        self,
3675        query: &[u8],
3676        db: &[u8],
3677        byte_len: usize,
3678        out: &mut [u32],
3679        index_of: impl Fn(usize) -> usize,
3680    ) {
3681        assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3682        if byte_len == 0 || out.is_empty() {
3683            return;
3684        }
3685        let row = |index: usize| -> &[u8] {
3686            let start = index * byte_len;
3687            &db[start..start + byte_len]
3688        };
3689        macro_rules! score_with {
3690            ($one:expr, $four:expr) => {{
3691                let mut i = 0;
3692                while i + HAMMING_ROWS_PER_KERNEL <= out.len() {
3693                    let quad = [
3694                        row(index_of(i)),
3695                        row(index_of(i + 1)),
3696                        row(index_of(i + 2)),
3697                        row(index_of(i + 3)),
3698                    ];
3699                    out[i..i + HAMMING_ROWS_PER_KERNEL].copy_from_slice(&$four(query, quad));
3700                    i += HAMMING_ROWS_PER_KERNEL;
3701                }
3702                while i < out.len() {
3703                    out[i] = $one(query, row(index_of(i)));
3704                    i += 1;
3705                }
3706            }};
3707        }
3708        match self {
3709            #[cfg(target_arch = "x86_64")]
3710            Self::Avx512 => score_with!(
3711                |query, row| unsafe { hamming_distance_avx512(query, row) },
3712                |query, rows| unsafe { hamming_distance_x4_avx512(query, rows) }
3713            ),
3714            #[cfg(target_arch = "x86_64")]
3715            Self::Avx2 => score_with!(
3716                |query, row| unsafe { avx2::hamming_distance(query, row) },
3717                |query, rows| unsafe { avx2::hamming_distance_x4(query, rows) }
3718            ),
3719            #[cfg(target_arch = "aarch64")]
3720            Self::Neon => score_with!(
3721                |query, row| unsafe { neon::hamming_distance(query, row) },
3722                |query, rows| unsafe { neon::hamming_distance_x4(query, rows) }
3723            ),
3724            Self::Scalar => score_with!(hamming_distance_scalar, hamming_distance_x4_scalar),
3725        }
3726    }
3727}
3728
3729/// Compute Hamming distance between two packed-bit vectors.
3730/// Returns the number of differing bits.
3731///
3732/// Uses NEON on aarch64 and VPOPCNTDQ/AVX2 on x86_64, with a scalar fallback.
3733/// Loops over many pairs should resolve a [`HammingKernel`] once instead of
3734/// paying feature detection here per pair.
3735#[inline]
3736pub fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
3737    assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3738    HammingKernel::resolve().distance(a, b)
3739}
3740
3741/// Scalar Hamming distance using u64 chunks + count_ones().
3742/// On x86_64, count_ones() compiles to POPCNT when target-cpu supports it.
3743#[inline]
3744#[allow(dead_code)]
3745fn hamming_distance_scalar(a: &[u8], b: &[u8]) -> u32 {
3746    let len = a.len();
3747    let chunks = len / 8;
3748    let remainder = len % 8;
3749    let mut total = 0u32;
3750
3751    for i in 0..chunks {
3752        let off = i * 8;
3753        let va = unsafe { std::ptr::read_unaligned(a.as_ptr().add(off) as *const u64) };
3754        let vb = unsafe { std::ptr::read_unaligned(b.as_ptr().add(off) as *const u64) };
3755        total += (va ^ vb).count_ones();
3756    }
3757
3758    let base = chunks * 8;
3759    for i in 0..remainder {
3760        total += (a[base + i] ^ b[base + i]).count_ones();
3761    }
3762
3763    total
3764}
3765
3766/// Batch Hamming scoring: compute similarity scores for multiple binary vectors.
3767///
3768/// `query` and each vector in `db` are packed-bit vectors of `byte_len` bytes each.
3769/// `dim_bits` is the number of bits (dimensions) for normalization.
3770/// Score = 1.0 - hamming_distance / dim_bits (range [0.0, 1.0]).
3771pub fn batch_hamming_scores(
3772    query: &[u8],
3773    db: &[u8],
3774    byte_len: usize,
3775    dim_bits: usize,
3776    scores: &mut [f32],
3777) {
3778    let n = scores.len();
3779    let required = n
3780        .checked_mul(byte_len)
3781        .expect("Hamming batch byte length overflow");
3782    assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3783    assert!(
3784        db.len() >= required,
3785        "Hamming batch is truncated: need {required} bytes, got {}",
3786        db.len()
3787    );
3788
3789    if byte_len == 0 || n == 0 || dim_bits == 0 {
3790        return;
3791    }
3792
3793    scores_from_hamming(
3794        HammingKernel::resolve(),
3795        query,
3796        db,
3797        byte_len,
3798        dim_bits,
3799        scores,
3800    );
3801}
3802
3803/// Batch Hamming scoring with a caller-resolved kernel.
3804///
3805/// Scans that already hold a [`HammingKernel`] (leaf scanning, Lloyd
3806/// assignment) use this to keep feature detection out of the loop entirely.
3807pub fn scores_from_hamming(
3808    kernel: HammingKernel,
3809    query: &[u8],
3810    db: &[u8],
3811    byte_len: usize,
3812    dim_bits: usize,
3813    scores: &mut [f32],
3814) {
3815    if byte_len == 0 || scores.is_empty() || dim_bits == 0 {
3816        return;
3817    }
3818    let inv_dim = 1.0 / dim_bits as f32;
3819    // Distances stay integral until the very last step; the stack block keeps
3820    // the row-batched kernel reachable without a per-scan allocation.
3821    let mut distances = [0u32; HAMMING_DISTANCE_BLOCK];
3822    for (block_index, block) in scores.chunks_mut(HAMMING_DISTANCE_BLOCK).enumerate() {
3823        let rows = &mut distances[..block.len()];
3824        kernel.distances(
3825            query,
3826            &db[block_index * HAMMING_DISTANCE_BLOCK * byte_len..],
3827            byte_len,
3828            rows,
3829        );
3830        for (score, &distance) in block.iter_mut().zip(rows.iter()) {
3831            *score = 1.0 - distance as f32 * inv_dim;
3832        }
3833    }
3834}
3835
3836/// Rows per stack block when converting batched distances into scores.
3837const HAMMING_DISTANCE_BLOCK: usize = 64;
3838
3839/// Batch Hamming distances (exact bit counts) for `out.len()` rows of `db`.
3840///
3841/// Callers that rank by distance — coarse assignment, routing — avoid the
3842/// float round-trip entirely.
3843pub fn batch_hamming_distances(query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3844    HammingKernel::resolve().distances(query, db, byte_len, out);
3845}
3846
3847#[cfg(test)]
3848mod tests {
3849    use super::*;
3850
3851    #[test]
3852    fn vector_simd_boundaries_reject_dimension_mismatches() {
3853        let vectors = vec![1.0f32; 6];
3854        let raw_f16 = vec![0u8; 12];
3855        let raw_u8 = vec![0u8; 6];
3856        let mut scores = vec![0.0f32; 2];
3857
3858        for invalid_query in [vec![1.0, 2.0], vec![1.0, 2.0, 3.0, 4.0]] {
3859            assert!(
3860                std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3861                    batch_cosine_scores(&invalid_query, &vectors, 3, &mut scores)
3862                }))
3863                .is_err()
3864            );
3865            assert!(
3866                std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3867                    batch_dot_scores_f16(&invalid_query, &raw_f16, 3, &mut scores)
3868                }))
3869                .is_err()
3870            );
3871            assert!(
3872                std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3873                    batch_cosine_scores_u8(&invalid_query, &raw_u8, 3, &mut scores)
3874                }))
3875                .is_err()
3876            );
3877        }
3878    }
3879
3880    #[test]
3881    fn vector_simd_boundaries_reject_truncated_storage() {
3882        let query = [1.0f32, 2.0, 3.0];
3883        let mut scores = [0.0f32; 2];
3884
3885        assert!(
3886            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3887                batch_dot_scores(&query, &[0.0; 5], 3, &mut scores)
3888            }))
3889            .is_err()
3890        );
3891        assert!(
3892            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3893                batch_cosine_scores_f16(&query, &[0u8; 11], 3, &mut scores)
3894            }))
3895            .is_err()
3896        );
3897        assert!(
3898            std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3899                dot_product_f32(&query, &query, 4)
3900            }))
3901            .is_err()
3902        );
3903    }
3904
3905    #[test]
3906    fn test_unpack_8bit() {
3907        let input: Vec<u8> = (0..128).collect();
3908        let mut output = vec![0u32; 128];
3909        unpack_8bit(&input, &mut output, 128);
3910
3911        for (i, &v) in output.iter().enumerate() {
3912            assert_eq!(v, i as u32);
3913        }
3914    }
3915
3916    #[test]
3917    fn test_unpack_16bit() {
3918        let mut input = vec![0u8; 256];
3919        for i in 0..128 {
3920            let val = (i * 100) as u16;
3921            input[i * 2] = val as u8;
3922            input[i * 2 + 1] = (val >> 8) as u8;
3923        }
3924
3925        let mut output = vec![0u32; 128];
3926        unpack_16bit(&input, &mut output, 128);
3927
3928        for (i, &v) in output.iter().enumerate() {
3929            assert_eq!(v, (i * 100) as u32);
3930        }
3931    }
3932
3933    #[test]
3934    fn test_unpack_32bit() {
3935        let mut input = vec![0u8; 512];
3936        for i in 0..128 {
3937            let val = (i * 1000) as u32;
3938            let bytes = val.to_le_bytes();
3939            input[i * 4..i * 4 + 4].copy_from_slice(&bytes);
3940        }
3941
3942        let mut output = vec![0u32; 128];
3943        unpack_32bit(&input, &mut output, 128);
3944
3945        for (i, &v) in output.iter().enumerate() {
3946            assert_eq!(v, (i * 1000) as u32);
3947        }
3948    }
3949
3950    #[test]
3951    fn test_delta_decode() {
3952        // doc_ids: [10, 15, 20, 30, 50]
3953        // gaps: [5, 5, 10, 20]
3954        // deltas (gap-1): [4, 4, 9, 19]
3955        let deltas = vec![4u32, 4, 9, 19];
3956        let mut output = vec![0u32; 5];
3957
3958        delta_decode(&mut output, &deltas, 10, 5);
3959
3960        assert_eq!(output, vec![10, 15, 20, 30, 50]);
3961    }
3962
3963    #[test]
3964    fn test_add_one() {
3965        let mut values = vec![0u32, 1, 2, 3, 4, 5, 6, 7];
3966        add_one(&mut values, 8);
3967
3968        assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
3969    }
3970
3971    #[test]
3972    fn test_bits_needed() {
3973        assert_eq!(bits_needed(0), 0);
3974        assert_eq!(bits_needed(1), 1);
3975        assert_eq!(bits_needed(2), 2);
3976        assert_eq!(bits_needed(3), 2);
3977        assert_eq!(bits_needed(4), 3);
3978        assert_eq!(bits_needed(255), 8);
3979        assert_eq!(bits_needed(256), 9);
3980        assert_eq!(bits_needed(u32::MAX), 32);
3981    }
3982
3983    #[test]
3984    fn test_unpack_8bit_delta_decode() {
3985        // doc_ids: [10, 15, 20, 30, 50]
3986        // gaps: [5, 5, 10, 20]
3987        // deltas (gap-1): [4, 4, 9, 19] stored as u8
3988        let input: Vec<u8> = vec![4, 4, 9, 19];
3989        let mut output = vec![0u32; 5];
3990
3991        unpack_8bit_delta_decode(&input, &mut output, 10, 5);
3992
3993        assert_eq!(output, vec![10, 15, 20, 30, 50]);
3994    }
3995
3996    #[test]
3997    fn test_unpack_16bit_delta_decode() {
3998        // doc_ids: [100, 600, 1100, 2100, 4100]
3999        // gaps: [500, 500, 1000, 2000]
4000        // deltas (gap-1): [499, 499, 999, 1999] stored as u16
4001        let mut input = vec![0u8; 8];
4002        for (i, &delta) in [499u16, 499, 999, 1999].iter().enumerate() {
4003            input[i * 2] = delta as u8;
4004            input[i * 2 + 1] = (delta >> 8) as u8;
4005        }
4006        let mut output = vec![0u32; 5];
4007
4008        unpack_16bit_delta_decode(&input, &mut output, 100, 5);
4009
4010        assert_eq!(output, vec![100, 600, 1100, 2100, 4100]);
4011    }
4012
4013    #[test]
4014    fn test_fused_vs_separate_8bit() {
4015        // Test that fused and separate operations produce the same result
4016        let input: Vec<u8> = (0..127).collect();
4017        let first_value = 1000u32;
4018        let count = 128;
4019
4020        // Separate: unpack then delta_decode
4021        let mut unpacked = vec![0u32; 128];
4022        unpack_8bit(&input, &mut unpacked, 127);
4023        let mut separate_output = vec![0u32; 128];
4024        delta_decode(&mut separate_output, &unpacked, first_value, count);
4025
4026        // Fused
4027        let mut fused_output = vec![0u32; 128];
4028        unpack_8bit_delta_decode(&input, &mut fused_output, first_value, count);
4029
4030        assert_eq!(separate_output, fused_output);
4031    }
4032
4033    #[test]
4034    fn test_round_bit_width() {
4035        assert_eq!(round_bit_width(0), 0);
4036        assert_eq!(round_bit_width(1), 8);
4037        assert_eq!(round_bit_width(5), 8);
4038        assert_eq!(round_bit_width(8), 8);
4039        assert_eq!(round_bit_width(9), 16);
4040        assert_eq!(round_bit_width(12), 16);
4041        assert_eq!(round_bit_width(16), 16);
4042        assert_eq!(round_bit_width(17), 32);
4043        assert_eq!(round_bit_width(24), 32);
4044        assert_eq!(round_bit_width(32), 32);
4045    }
4046
4047    #[test]
4048    fn test_rounded_bitwidth_from_exact() {
4049        assert_eq!(RoundedBitWidth::from_exact(0), RoundedBitWidth::Zero);
4050        assert_eq!(RoundedBitWidth::from_exact(1), RoundedBitWidth::Bits8);
4051        assert_eq!(RoundedBitWidth::from_exact(8), RoundedBitWidth::Bits8);
4052        assert_eq!(RoundedBitWidth::from_exact(9), RoundedBitWidth::Bits16);
4053        assert_eq!(RoundedBitWidth::from_exact(16), RoundedBitWidth::Bits16);
4054        assert_eq!(RoundedBitWidth::from_exact(17), RoundedBitWidth::Bits32);
4055        assert_eq!(RoundedBitWidth::from_exact(32), RoundedBitWidth::Bits32);
4056    }
4057
4058    #[test]
4059    fn test_pack_unpack_rounded_8bit() {
4060        let values: Vec<u32> = (0..128).map(|i| i % 256).collect();
4061        let mut packed = vec![0u8; 128];
4062
4063        let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits8, &mut packed);
4064        assert_eq!(bytes_written, 128);
4065
4066        let mut unpacked = vec![0u32; 128];
4067        unpack_rounded(&packed, RoundedBitWidth::Bits8, &mut unpacked, 128);
4068
4069        assert_eq!(values, unpacked);
4070    }
4071
4072    #[test]
4073    fn test_pack_unpack_rounded_16bit() {
4074        let values: Vec<u32> = (0..128).map(|i| i * 100).collect();
4075        let mut packed = vec![0u8; 256];
4076
4077        let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits16, &mut packed);
4078        assert_eq!(bytes_written, 256);
4079
4080        let mut unpacked = vec![0u32; 128];
4081        unpack_rounded(&packed, RoundedBitWidth::Bits16, &mut unpacked, 128);
4082
4083        assert_eq!(values, unpacked);
4084    }
4085
4086    #[test]
4087    fn test_pack_unpack_rounded_32bit() {
4088        let values: Vec<u32> = (0..128).map(|i| i * 100000).collect();
4089        let mut packed = vec![0u8; 512];
4090
4091        let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits32, &mut packed);
4092        assert_eq!(bytes_written, 512);
4093
4094        let mut unpacked = vec![0u32; 128];
4095        unpack_rounded(&packed, RoundedBitWidth::Bits32, &mut unpacked, 128);
4096
4097        assert_eq!(values, unpacked);
4098    }
4099
4100    #[test]
4101    fn test_unpack_rounded_delta_decode() {
4102        // Test 8-bit rounded delta decode
4103        // doc_ids: [10, 15, 20, 30, 50]
4104        // gaps: [5, 5, 10, 20]
4105        // deltas (gap-1): [4, 4, 9, 19] stored as u8
4106        let input: Vec<u8> = vec![4, 4, 9, 19];
4107        let mut output = vec![0u32; 5];
4108
4109        unpack_rounded_delta_decode(&input, RoundedBitWidth::Bits8, &mut output, 10, 5);
4110
4111        assert_eq!(output, vec![10, 15, 20, 30, 50]);
4112    }
4113
4114    #[test]
4115    fn test_unpack_rounded_delta_decode_zero() {
4116        // All zeros means gaps of 1 (consecutive doc IDs)
4117        let input: Vec<u8> = vec![];
4118        let mut output = vec![0u32; 5];
4119
4120        unpack_rounded_delta_decode(&input, RoundedBitWidth::Zero, &mut output, 100, 5);
4121
4122        assert_eq!(output, vec![100, 101, 102, 103, 104]);
4123    }
4124
4125    // ========================================================================
4126    // Sparse Vector SIMD Tests
4127    // ========================================================================
4128
4129    #[test]
4130    fn test_dequantize_uint8() {
4131        let input: Vec<u8> = vec![0, 128, 255, 64, 192];
4132        let mut output = vec![0.0f32; 5];
4133        let scale = 0.1;
4134        let min_val = 1.0;
4135
4136        dequantize_uint8(&input, &mut output, scale, min_val, 5);
4137
4138        // Expected: input[i] * scale + min_val
4139        assert!((output[0] - 1.0).abs() < 1e-6); // 0 * 0.1 + 1.0 = 1.0
4140        assert!((output[1] - 13.8).abs() < 1e-6); // 128 * 0.1 + 1.0 = 13.8
4141        assert!((output[2] - 26.5).abs() < 1e-6); // 255 * 0.1 + 1.0 = 26.5
4142        assert!((output[3] - 7.4).abs() < 1e-6); // 64 * 0.1 + 1.0 = 7.4
4143        assert!((output[4] - 20.2).abs() < 1e-6); // 192 * 0.1 + 1.0 = 20.2
4144    }
4145
4146    #[test]
4147    fn test_dequantize_uint8_large() {
4148        // Test with 128 values (full SIMD block)
4149        let input: Vec<u8> = (0..128).collect();
4150        let mut output = vec![0.0f32; 128];
4151        let scale = 2.0;
4152        let min_val = -10.0;
4153
4154        dequantize_uint8(&input, &mut output, scale, min_val, 128);
4155
4156        for (i, &out) in output.iter().enumerate().take(128) {
4157            let expected = i as f32 * scale + min_val;
4158            assert!(
4159                (out - expected).abs() < 1e-5,
4160                "Mismatch at {}: expected {}, got {}",
4161                i,
4162                expected,
4163                out
4164            );
4165        }
4166    }
4167
4168    #[test]
4169    fn test_dot_product_f32() {
4170        let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
4171        let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0];
4172
4173        let result = dot_product_f32(&a, &b, 5);
4174
4175        // Expected: 1*2 + 2*3 + 3*4 + 4*5 + 5*6 = 2 + 6 + 12 + 20 + 30 = 70
4176        assert!((result - 70.0).abs() < 1e-5);
4177    }
4178
4179    #[test]
4180    fn test_dot_product_f32_large() {
4181        // Test with 128 values
4182        let a: Vec<f32> = (0..128).map(|i| i as f32).collect();
4183        let b: Vec<f32> = (0..128).map(|i| (i + 1) as f32).collect();
4184
4185        let result = dot_product_f32(&a, &b, 128);
4186
4187        // Compute expected
4188        let expected: f32 = (0..128).map(|i| (i as f32) * ((i + 1) as f32)).sum();
4189        assert!(
4190            (result - expected).abs() < 1e-3,
4191            "Expected {}, got {}",
4192            expected,
4193            result
4194        );
4195    }
4196
4197    #[test]
4198    fn test_fused_dot_norm() {
4199        let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
4200        let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
4201        let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4202
4203        let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4204        let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4205        assert!(
4206            (dot - expected_dot).abs() < 1e-5,
4207            "dot: expected {}, got {}",
4208            expected_dot,
4209            dot
4210        );
4211        assert!(
4212            (norm_b - expected_norm).abs() < 1e-5,
4213            "norm: expected {}, got {}",
4214            expected_norm,
4215            norm_b
4216        );
4217    }
4218
4219    #[test]
4220    fn test_fused_dot_norm_large() {
4221        let a: Vec<f32> = (0..768).map(|i| (i as f32) * 0.01).collect();
4222        let b: Vec<f32> = (0..768).map(|i| (i as f32) * 0.02 + 0.5).collect();
4223        let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4224
4225        let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4226        let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4227        assert!(
4228            (dot - expected_dot).abs() < 1.0,
4229            "dot: expected {}, got {}",
4230            expected_dot,
4231            dot
4232        );
4233        assert!(
4234            (norm_b - expected_norm).abs() < 1.0,
4235            "norm: expected {}, got {}",
4236            expected_norm,
4237            norm_b
4238        );
4239    }
4240
4241    #[test]
4242    fn test_batch_cosine_scores() {
4243        // 4 vectors of dim 3
4244        let query = vec![1.0f32, 0.0, 0.0];
4245        let vectors = vec![
4246            1.0, 0.0, 0.0, // identical to query
4247            0.0, 1.0, 0.0, // orthogonal
4248            -1.0, 0.0, 0.0, // opposite
4249            0.5, 0.5, 0.0, // 45 degrees
4250        ];
4251        let mut scores = vec![0f32; 4];
4252        batch_cosine_scores(&query, &vectors, 3, &mut scores);
4253
4254        assert!((scores[0] - 1.0).abs() < 1e-5, "identical: {}", scores[0]);
4255        assert!(scores[1].abs() < 1e-5, "orthogonal: {}", scores[1]);
4256        assert!((scores[2] - (-1.0)).abs() < 1e-5, "opposite: {}", scores[2]);
4257        let expected_45 = 0.5f32 / (0.5f32.powi(2) + 0.5f32.powi(2)).sqrt();
4258        assert!(
4259            (scores[3] - expected_45).abs() < 1e-5,
4260            "45deg: expected {}, got {}",
4261            expected_45,
4262            scores[3]
4263        );
4264    }
4265
4266    #[test]
4267    fn test_batch_cosine_scores_matches_individual() {
4268        let query: Vec<f32> = (0..128).map(|i| (i as f32) * 0.1).collect();
4269        let n = 50;
4270        let dim = 128;
4271        let vectors: Vec<f32> = (0..n * dim).map(|i| ((i * 7 + 3) as f32) * 0.01).collect();
4272
4273        let mut batch_scores = vec![0f32; n];
4274        batch_cosine_scores(&query, &vectors, dim, &mut batch_scores);
4275
4276        for i in 0..n {
4277            let vec_i = &vectors[i * dim..(i + 1) * dim];
4278            let individual = cosine_similarity(&query, vec_i);
4279            assert!(
4280                (batch_scores[i] - individual).abs() < 1e-5,
4281                "vec {}: batch={}, individual={}",
4282                i,
4283                batch_scores[i],
4284                individual
4285            );
4286        }
4287    }
4288
4289    #[test]
4290    fn test_batch_cosine_scores_empty() {
4291        let query = vec![1.0f32, 2.0, 3.0];
4292        let vectors: Vec<f32> = vec![];
4293        let mut scores: Vec<f32> = vec![];
4294        batch_cosine_scores(&query, &vectors, 3, &mut scores);
4295        assert!(scores.is_empty());
4296    }
4297
4298    #[test]
4299    fn test_batch_cosine_scores_zero_query() {
4300        let query = vec![0.0f32, 0.0, 0.0];
4301        let vectors = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
4302        let mut scores = vec![0f32; 2];
4303        batch_cosine_scores(&query, &vectors, 3, &mut scores);
4304        assert_eq!(scores[0], 0.0);
4305        assert_eq!(scores[1], 0.0);
4306    }
4307
4308    // ================================================================
4309    // f16 conversion tests
4310    // ================================================================
4311
4312    #[test]
4313    fn test_f16_roundtrip_normal() {
4314        for &v in &[0.0f32, 1.0, -1.0, 0.5, -0.5, 0.333, 65504.0] {
4315            let h = f32_to_f16(v);
4316            let back = f16_to_f32(h);
4317            let err = (back - v).abs() / v.abs().max(1e-6);
4318            assert!(
4319                err < 0.002,
4320                "f16 roundtrip {v} → {h:#06x} → {back}, rel err {err}"
4321            );
4322        }
4323    }
4324
4325    #[test]
4326    fn test_f16_special() {
4327        // Zero
4328        assert_eq!(f16_to_f32(f32_to_f16(0.0)), 0.0);
4329        // Negative zero
4330        assert_eq!(f32_to_f16(-0.0), 0x8000);
4331        // Infinity
4332        assert!(f16_to_f32(f32_to_f16(f32::INFINITY)).is_infinite());
4333        // NaN
4334        assert!(f16_to_f32(f32_to_f16(f32::NAN)).is_nan());
4335    }
4336
4337    #[test]
4338    fn test_f16_embedding_range() {
4339        // Typical embedding values in [-1, 1]
4340        let values: Vec<f32> = (-100..=100).map(|i| i as f32 / 100.0).collect();
4341        for &v in &values {
4342            let back = f16_to_f32(f32_to_f16(v));
4343            assert!((back - v).abs() < 0.001, "f16 error for {v}: got {back}");
4344        }
4345    }
4346
4347    // ================================================================
4348    // u8 conversion tests
4349    // ================================================================
4350
4351    #[test]
4352    fn test_u8_roundtrip() {
4353        // Boundary values
4354        assert_eq!(f32_to_u8_saturating(-1.0), 0);
4355        assert_eq!(f32_to_u8_saturating(1.0), 255);
4356        assert_eq!(f32_to_u8_saturating(0.0), 127); // ~127.5 truncated
4357
4358        // Saturation
4359        assert_eq!(f32_to_u8_saturating(-2.0), 0);
4360        assert_eq!(f32_to_u8_saturating(2.0), 255);
4361    }
4362
4363    #[test]
4364    fn test_u8_dequantize() {
4365        assert!((u8_to_f32(0) - (-1.0)).abs() < 0.01);
4366        assert!((u8_to_f32(255) - 1.0).abs() < 0.01);
4367        assert!((u8_to_f32(127) - 0.0).abs() < 0.01);
4368    }
4369
4370    // ================================================================
4371    // Batch scoring tests for quantized vectors
4372    // ================================================================
4373
4374    #[test]
4375    fn test_batch_cosine_scores_f16() {
4376        let query = vec![0.6f32, 0.8, 0.0, 0.0];
4377        let dim = 4;
4378        let vecs_f32 = vec![
4379            0.6f32, 0.8, 0.0, 0.0, // identical to query
4380            0.0, 0.0, 0.6, 0.8, // orthogonal
4381        ];
4382
4383        // Quantize to f16
4384        let mut f16_buf = vec![0u16; 8];
4385        batch_f32_to_f16(&vecs_f32, &mut f16_buf);
4386        let raw: &[u8] =
4387            unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4388
4389        let mut scores = vec![0f32; 2];
4390        batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4391
4392        assert!(
4393            (scores[0] - 1.0).abs() < 0.01,
4394            "identical vectors: {}",
4395            scores[0]
4396        );
4397        assert!(scores[1].abs() < 0.01, "orthogonal vectors: {}", scores[1]);
4398    }
4399
4400    #[test]
4401    fn test_batch_cosine_scores_u8() {
4402        let query = vec![0.6f32, 0.8, 0.0, 0.0];
4403        let dim = 4;
4404        let vecs_f32 = vec![
4405            0.6f32, 0.8, 0.0, 0.0, // ~identical to query
4406            -0.6, -0.8, 0.0, 0.0, // opposite
4407        ];
4408
4409        // Quantize to u8
4410        let mut u8_buf = vec![0u8; 8];
4411        batch_f32_to_u8(&vecs_f32, &mut u8_buf);
4412
4413        let mut scores = vec![0f32; 2];
4414        batch_cosine_scores_u8(&query, &u8_buf, dim, &mut scores);
4415
4416        assert!(scores[0] > 0.95, "similar vectors: {}", scores[0]);
4417        assert!(scores[1] < -0.95, "opposite vectors: {}", scores[1]);
4418    }
4419
4420    #[test]
4421    fn test_batch_cosine_scores_f16_large_dim() {
4422        // Test with typical embedding dimension
4423        let dim = 768;
4424        let query: Vec<f32> = (0..dim).map(|i| (i as f32 / dim as f32) - 0.5).collect();
4425        let vec2: Vec<f32> = query.iter().map(|x| x * 0.9 + 0.01).collect();
4426
4427        let mut all_vecs = query.clone();
4428        all_vecs.extend_from_slice(&vec2);
4429
4430        let mut f16_buf = vec![0u16; all_vecs.len()];
4431        batch_f32_to_f16(&all_vecs, &mut f16_buf);
4432        let raw: &[u8] =
4433            unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4434
4435        let mut scores = vec![0f32; 2];
4436        batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4437
4438        // Self-similarity should be ~1.0
4439        assert!((scores[0] - 1.0).abs() < 0.01, "self-sim: {}", scores[0]);
4440        // High similarity with scaled version
4441        assert!(scores[1] > 0.99, "scaled-sim: {}", scores[1]);
4442    }
4443
4444    // ================================================================
4445    // Hamming distance tests
4446    // ================================================================
4447
4448    #[test]
4449    fn test_hamming_distance_identical() {
4450        let a = vec![0xAA; 64];
4451        assert_eq!(hamming_distance(&a, &a), 0);
4452    }
4453
4454    #[test]
4455    fn test_hamming_distance_opposite() {
4456        let a = vec![0xFF; 32];
4457        let b = vec![0x00; 32];
4458        assert_eq!(hamming_distance(&a, &b), 256);
4459    }
4460
4461    #[test]
4462    fn test_hamming_distance_known() {
4463        // Single byte: 0b10101010 vs 0b01010101 = 8 bits differ
4464        let a = vec![0xAA];
4465        let b = vec![0x55];
4466        assert_eq!(hamming_distance(&a, &b), 8);
4467
4468        // Two bytes
4469        let a = vec![0xFF, 0x00];
4470        let b = vec![0x00, 0x00];
4471        assert_eq!(hamming_distance(&a, &b), 8);
4472    }
4473
4474    #[test]
4475    fn test_hamming_distance_single_bit() {
4476        let a = vec![0x00; 16];
4477        let mut b = vec![0x00; 16];
4478        b[7] = 0x01; // flip one bit
4479        assert_eq!(hamming_distance(&a, &b), 1);
4480    }
4481
4482    #[test]
4483    fn test_hamming_distance_empty() {
4484        let a: Vec<u8> = vec![];
4485        assert_eq!(hamming_distance(&a, &a), 0);
4486    }
4487
4488    #[test]
4489    fn test_hamming_distance_remainder_path() {
4490        // 17 bytes: not aligned to 16 (NEON) or 32 (AVX2)
4491        let a = vec![0xFF; 17];
4492        let b = vec![0x00; 17];
4493        assert_eq!(hamming_distance(&a, &b), 136); // 17 * 8
4494
4495        // 33 bytes: tests 32-byte chunk + 1 remainder for AVX2
4496        let a = vec![0xFF; 33];
4497        let b = vec![0x00; 33];
4498        assert_eq!(hamming_distance(&a, &b), 264); // 33 * 8
4499    }
4500
4501    #[test]
4502    fn test_hamming_distance_large() {
4503        // 4096 bytes = 32768 bits, all differing
4504        let a = vec![0xFF; 4096];
4505        let b = vec![0x00; 4096];
4506        assert_eq!(hamming_distance(&a, &b), 32768);
4507    }
4508
4509    #[test]
4510    fn test_hamming_distance_scalar_matches() {
4511        // Verify SIMD path matches scalar for various sizes
4512        for size in [1, 7, 8, 15, 16, 31, 32, 63, 64, 100, 128, 255, 256] {
4513            let a: Vec<u8> = (0..size).map(|i| (i * 37 + 13) as u8).collect();
4514            let b: Vec<u8> = (0..size).map(|i| (i * 53 + 7) as u8).collect();
4515            let expected = hamming_distance_scalar(&a, &b);
4516            let got = hamming_distance(&a, &b);
4517            assert_eq!(got, expected, "mismatch at size {size}");
4518        }
4519    }
4520
4521    // ================================================================
4522    // Batch Hamming scoring tests
4523    // ================================================================
4524
4525    #[test]
4526    fn test_batch_hamming_scores_identical() {
4527        let query = vec![0xAA; 16];
4528        let db = vec![0xAA; 16]; // one vector, identical
4529        let mut scores = vec![0f32; 1];
4530        batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4531        assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4532    }
4533
4534    #[test]
4535    fn test_batch_hamming_scores_opposite() {
4536        let query = vec![0xFF; 16];
4537        let db = vec![0x00; 16];
4538        let mut scores = vec![0f32; 1];
4539        batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4540        assert!((scores[0] - 0.0).abs() < 1e-6, "opposite: {}", scores[0]);
4541    }
4542
4543    #[test]
4544    fn test_batch_hamming_scores_multiple() {
4545        let byte_len = 8;
4546        let dim_bits = 64;
4547        let query = vec![0xFF; byte_len];
4548        let mut db = Vec::new();
4549        db.extend_from_slice(&vec![0xFF; byte_len]); // identical → 1.0
4550        db.extend_from_slice(&vec![0x00; byte_len]); // opposite → 0.0
4551        db.extend_from_slice(&vec![0x0F; byte_len]); // half bits differ → 0.5
4552
4553        let mut scores = vec![0f32; 3];
4554        batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut scores);
4555
4556        assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4557        assert!((scores[1] - 0.0).abs() < 1e-6, "opposite: {}", scores[1]);
4558        assert!((scores[2] - 0.5).abs() < 1e-6, "half: {}", scores[2]);
4559    }
4560
4561    #[test]
4562    fn test_batch_hamming_scores_empty() {
4563        let query = vec![0xFF; 8];
4564        let db: Vec<u8> = vec![];
4565        let mut scores: Vec<f32> = vec![];
4566        batch_hamming_scores(&query, &db, 8, 64, &mut scores);
4567        assert!(scores.is_empty());
4568    }
4569
4570    #[test]
4571    fn test_batch_hamming_scores_zero_byte_len() {
4572        let query: Vec<u8> = vec![];
4573        let db: Vec<u8> = vec![];
4574        let mut scores = vec![0f32; 1];
4575        batch_hamming_scores(&query, &db, 0, 0, &mut scores);
4576        // Should return early without modifying scores
4577        assert_eq!(scores[0], 0.0);
4578    }
4579
4580    // ================================================================
4581    // Resolved-kernel and row-batched Hamming tests
4582    // ================================================================
4583
4584    fn hamming_matrix(rows: usize, byte_len: usize) -> (Vec<u8>, Vec<u8>) {
4585        let query: Vec<u8> = (0..byte_len).map(|i| (i * 31 + 5) as u8).collect();
4586        let db: Vec<u8> = (0..rows * byte_len)
4587            .map(|i| (i * 97 + i / byte_len * 11 + 3) as u8)
4588            .collect();
4589        (query, db)
4590    }
4591
4592    /// The row-batched kernels share query loads and accumulators across four
4593    /// rows; every width must still agree bit-for-bit with the scalar loop.
4594    #[test]
4595    fn batched_hamming_distances_match_scalar_for_every_row_count() {
4596        let kernel = HammingKernel::resolve();
4597        // Cover both multiples of the quad width and every tail remainder, and
4598        // byte lengths that exercise 16/32/64-byte chunking plus odd tails.
4599        for byte_len in [1, 7, 8, 15, 16, 31, 32, 33, 63, 64, 65, 128, 320] {
4600            for rows in [1, 2, 3, 4, 5, 7, 8, 9, 64, 70] {
4601                let (query, db) = hamming_matrix(rows, byte_len);
4602                let mut got = vec![0u32; rows];
4603                kernel.distances(&query, &db, byte_len, &mut got);
4604                for (row, &distance) in got.iter().enumerate() {
4605                    let expected =
4606                        hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4607                    assert_eq!(
4608                        distance, expected,
4609                        "row {row} of {rows} at byte_len {byte_len}"
4610                    );
4611                }
4612            }
4613        }
4614    }
4615
4616    #[test]
4617    fn gathered_hamming_distances_follow_row_ids() {
4618        let kernel = HammingKernel::resolve();
4619        let byte_len = 320;
4620        let rows = 37;
4621        let (query, db) = hamming_matrix(rows, byte_len);
4622        // Scattered, repeated and reversed ids: routing visits rows in graph
4623        // order, not storage order.
4624        let ids: Vec<u32> = [36, 0, 17, 17, 5, 31, 2, 9, 9, 36, 1].into_iter().collect();
4625        let mut got = vec![0u32; ids.len()];
4626        kernel.gather_distances(&query, &db, byte_len, &ids, &mut got);
4627        for (slot, &id) in ids.iter().enumerate() {
4628            let start = id as usize * byte_len;
4629            let expected = hamming_distance_scalar(&query, &db[start..start + byte_len]);
4630            assert_eq!(got[slot], expected, "slot {slot} for row {id}");
4631        }
4632    }
4633
4634    #[test]
4635    fn resolved_kernel_matches_scalar_pairwise() {
4636        let kernel = HammingKernel::resolve();
4637        for byte_len in [1, 8, 32, 64, 65, 320, 4096] {
4638            let (query, db) = hamming_matrix(1, byte_len);
4639            assert_eq!(
4640                kernel.distance(&query, &db),
4641                hamming_distance_scalar(&query, &db),
4642                "byte_len {byte_len}"
4643            );
4644        }
4645    }
4646
4647    #[test]
4648    fn scores_from_hamming_matches_batch_scores_across_blocks() {
4649        let kernel = HammingKernel::resolve();
4650        let byte_len = 320;
4651        let dim_bits = byte_len * 8;
4652        // More rows than one stack block so the block seam is covered.
4653        let rows = HAMMING_DISTANCE_BLOCK * 2 + 3;
4654        let (query, db) = hamming_matrix(rows, byte_len);
4655        let mut expected = vec![0f32; rows];
4656        for (row, score) in expected.iter_mut().enumerate() {
4657            let distance =
4658                hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4659            *score = 1.0 - distance as f32 / dim_bits as f32;
4660        }
4661        let mut got = vec![0f32; rows];
4662        scores_from_hamming(kernel, &query, &db, byte_len, dim_bits, &mut got);
4663        for (row, (&got, &want)) in got.iter().zip(expected.iter()).enumerate() {
4664            assert!((got - want).abs() < 1e-6, "row {row}: {got} vs {want}");
4665        }
4666        let mut public = vec![0f32; rows];
4667        batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut public);
4668        assert_eq!(got, public);
4669    }
4670}
4671
4672// ============================================================================
4673// SIMD-accelerated linear scan for sorted u32 slices (within-block seek)
4674// ============================================================================
4675
4676/// Find index of first element >= `target` in a sorted `u32` slice.
4677///
4678/// Equivalent to `slice.partition_point(|&d| d < target)` but uses SIMD to
4679/// scan 4 elements per cycle. Faster than binary search for slices ≤ 256
4680/// elements because it avoids the data-dependency chain inherent in binary
4681/// search (~8-10 cycles/iteration vs ~1-2 cycles/iteration for SIMD scan).
4682///
4683/// Returns `slice.len()` if no element >= `target`.
4684#[inline]
4685pub fn find_first_ge_u32(slice: &[u32], target: u32) -> usize {
4686    #[cfg(target_arch = "aarch64")]
4687    {
4688        if neon::is_available() {
4689            return unsafe { find_first_ge_u32_neon(slice, target) };
4690        }
4691    }
4692
4693    #[cfg(target_arch = "x86_64")]
4694    {
4695        if sse::is_available() {
4696            return unsafe { find_first_ge_u32_sse(slice, target) };
4697        }
4698    }
4699
4700    // Scalar fallback (WASM, other architectures)
4701    slice.partition_point(|&d| d < target)
4702}
4703
4704#[cfg(target_arch = "aarch64")]
4705#[target_feature(enable = "neon")]
4706#[allow(unsafe_op_in_unsafe_fn)]
4707unsafe fn find_first_ge_u32_neon(slice: &[u32], target: u32) -> usize {
4708    use std::arch::aarch64::*;
4709
4710    let n = slice.len();
4711    let ptr = slice.as_ptr();
4712    let target_vec = vdupq_n_u32(target);
4713    // Bit positions for each lane: [1, 2, 4, 8]
4714    let bit_mask: uint32x4_t = core::mem::transmute([1u32, 2u32, 4u32, 8u32]);
4715
4716    let chunks = n / 16;
4717    let mut base = 0usize;
4718
4719    // Process 16 elements per iteration (4 × 4-wide NEON compares)
4720    for _ in 0..chunks {
4721        let v0 = vld1q_u32(ptr.add(base));
4722        let v1 = vld1q_u32(ptr.add(base + 4));
4723        let v2 = vld1q_u32(ptr.add(base + 8));
4724        let v3 = vld1q_u32(ptr.add(base + 12));
4725
4726        let c0 = vcgeq_u32(v0, target_vec);
4727        let c1 = vcgeq_u32(v1, target_vec);
4728        let c2 = vcgeq_u32(v2, target_vec);
4729        let c3 = vcgeq_u32(v3, target_vec);
4730
4731        let m0 = vaddvq_u32(vandq_u32(c0, bit_mask));
4732        if m0 != 0 {
4733            return base + m0.trailing_zeros() as usize;
4734        }
4735        let m1 = vaddvq_u32(vandq_u32(c1, bit_mask));
4736        if m1 != 0 {
4737            return base + 4 + m1.trailing_zeros() as usize;
4738        }
4739        let m2 = vaddvq_u32(vandq_u32(c2, bit_mask));
4740        if m2 != 0 {
4741            return base + 8 + m2.trailing_zeros() as usize;
4742        }
4743        let m3 = vaddvq_u32(vandq_u32(c3, bit_mask));
4744        if m3 != 0 {
4745            return base + 12 + m3.trailing_zeros() as usize;
4746        }
4747        base += 16;
4748    }
4749
4750    // Process remaining 4 elements at a time
4751    while base + 4 <= n {
4752        let vals = vld1q_u32(ptr.add(base));
4753        let cmp = vcgeq_u32(vals, target_vec);
4754        let mask = vaddvq_u32(vandq_u32(cmp, bit_mask));
4755        if mask != 0 {
4756            return base + mask.trailing_zeros() as usize;
4757        }
4758        base += 4;
4759    }
4760
4761    // Scalar remainder (0-3 elements)
4762    while base < n {
4763        if *slice.get_unchecked(base) >= target {
4764            return base;
4765        }
4766        base += 1;
4767    }
4768    n
4769}
4770
4771#[cfg(target_arch = "x86_64")]
4772#[target_feature(enable = "sse2")]
4773#[allow(unsafe_op_in_unsafe_fn)]
4774unsafe fn find_first_ge_u32_sse(slice: &[u32], target: u32) -> usize {
4775    use std::arch::x86_64::*;
4776
4777    let n = slice.len();
4778    let ptr = slice.as_ptr();
4779
4780    // For unsigned >= comparison: XOR with 0x80000000 converts to signed domain
4781    let sign_flip = _mm_set1_epi32(i32::MIN);
4782    let target_xor = _mm_xor_si128(_mm_set1_epi32(target as i32), sign_flip);
4783
4784    let chunks = n / 16;
4785    let mut base = 0usize;
4786
4787    // Process 16 elements per iteration (4 × 4-wide SSE compares)
4788    for _ in 0..chunks {
4789        let v0 = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4790        let v1 = _mm_xor_si128(
4791            _mm_loadu_si128(ptr.add(base + 4) as *const __m128i),
4792            sign_flip,
4793        );
4794        let v2 = _mm_xor_si128(
4795            _mm_loadu_si128(ptr.add(base + 8) as *const __m128i),
4796            sign_flip,
4797        );
4798        let v3 = _mm_xor_si128(
4799            _mm_loadu_si128(ptr.add(base + 12) as *const __m128i),
4800            sign_flip,
4801        );
4802
4803        // ge = eq | gt (in signed domain after XOR)
4804        let ge0 = _mm_or_si128(
4805            _mm_cmpeq_epi32(v0, target_xor),
4806            _mm_cmpgt_epi32(v0, target_xor),
4807        );
4808        let m0 = _mm_movemask_ps(_mm_castsi128_ps(ge0)) as u32;
4809        if m0 != 0 {
4810            return base + m0.trailing_zeros() as usize;
4811        }
4812
4813        let ge1 = _mm_or_si128(
4814            _mm_cmpeq_epi32(v1, target_xor),
4815            _mm_cmpgt_epi32(v1, target_xor),
4816        );
4817        let m1 = _mm_movemask_ps(_mm_castsi128_ps(ge1)) as u32;
4818        if m1 != 0 {
4819            return base + 4 + m1.trailing_zeros() as usize;
4820        }
4821
4822        let ge2 = _mm_or_si128(
4823            _mm_cmpeq_epi32(v2, target_xor),
4824            _mm_cmpgt_epi32(v2, target_xor),
4825        );
4826        let m2 = _mm_movemask_ps(_mm_castsi128_ps(ge2)) as u32;
4827        if m2 != 0 {
4828            return base + 8 + m2.trailing_zeros() as usize;
4829        }
4830
4831        let ge3 = _mm_or_si128(
4832            _mm_cmpeq_epi32(v3, target_xor),
4833            _mm_cmpgt_epi32(v3, target_xor),
4834        );
4835        let m3 = _mm_movemask_ps(_mm_castsi128_ps(ge3)) as u32;
4836        if m3 != 0 {
4837            return base + 12 + m3.trailing_zeros() as usize;
4838        }
4839        base += 16;
4840    }
4841
4842    // Process remaining 4 elements at a time
4843    while base + 4 <= n {
4844        let vals = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4845        let ge = _mm_or_si128(
4846            _mm_cmpeq_epi32(vals, target_xor),
4847            _mm_cmpgt_epi32(vals, target_xor),
4848        );
4849        let mask = _mm_movemask_ps(_mm_castsi128_ps(ge)) as u32;
4850        if mask != 0 {
4851            return base + mask.trailing_zeros() as usize;
4852        }
4853        base += 4;
4854    }
4855
4856    // Scalar remainder (0-3 elements)
4857    while base < n {
4858        if *slice.get_unchecked(base) >= target {
4859            return base;
4860        }
4861        base += 1;
4862    }
4863    n
4864}
4865
4866#[cfg(test)]
4867mod find_first_ge_tests {
4868    use super::find_first_ge_u32;
4869
4870    #[test]
4871    fn test_find_first_ge_basic() {
4872        let data: Vec<u32> = (0..128).map(|i| i * 3).collect(); // [0, 3, 6, ..., 381]
4873        assert_eq!(find_first_ge_u32(&data, 0), 0);
4874        assert_eq!(find_first_ge_u32(&data, 1), 1); // first >= 1 is 3 at idx 1
4875        assert_eq!(find_first_ge_u32(&data, 3), 1);
4876        assert_eq!(find_first_ge_u32(&data, 4), 2); // first >= 4 is 6 at idx 2
4877        assert_eq!(find_first_ge_u32(&data, 381), 127);
4878        assert_eq!(find_first_ge_u32(&data, 382), 128); // past end
4879    }
4880
4881    #[test]
4882    fn test_find_first_ge_matches_partition_point() {
4883        let data: Vec<u32> = vec![1, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55, 60, 65, 70, 75];
4884        for target in 0..80 {
4885            let expected = data.partition_point(|&d| d < target);
4886            let actual = find_first_ge_u32(&data, target);
4887            assert_eq!(actual, expected, "target={}", target);
4888        }
4889    }
4890
4891    #[test]
4892    fn test_find_first_ge_small_slices() {
4893        // Empty
4894        assert_eq!(find_first_ge_u32(&[], 5), 0);
4895        // Single element
4896        assert_eq!(find_first_ge_u32(&[10], 5), 0);
4897        assert_eq!(find_first_ge_u32(&[10], 10), 0);
4898        assert_eq!(find_first_ge_u32(&[10], 11), 1);
4899        // Three elements (< SIMD width)
4900        assert_eq!(find_first_ge_u32(&[2, 4, 6], 5), 2);
4901    }
4902
4903    #[test]
4904    fn test_find_first_ge_full_block() {
4905        // Simulate a full 128-entry block
4906        let data: Vec<u32> = (100..228).collect();
4907        assert_eq!(find_first_ge_u32(&data, 100), 0);
4908        assert_eq!(find_first_ge_u32(&data, 150), 50);
4909        assert_eq!(find_first_ge_u32(&data, 227), 127);
4910        assert_eq!(find_first_ge_u32(&data, 228), 128);
4911        assert_eq!(find_first_ge_u32(&data, 99), 0);
4912    }
4913
4914    #[test]
4915    fn test_find_first_ge_u32_max() {
4916        // Test with large u32 values (unsigned correctness)
4917        let data = vec![u32::MAX - 10, u32::MAX - 5, u32::MAX - 1, u32::MAX];
4918        assert_eq!(find_first_ge_u32(&data, u32::MAX - 10), 0);
4919        assert_eq!(find_first_ge_u32(&data, u32::MAX - 7), 1);
4920        assert_eq!(find_first_ge_u32(&data, u32::MAX), 3);
4921    }
4922}