Skip to main content

summa_core/structures/
simd.rs

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