Skip to main content

lattice_embed/simd/
quantized.rs

1//! INT8 vector quantization and approximate similarity kernels.
2//!
3//! Constructor-owned values preserve the SIMD range invariant.
4//!
5//! See docs/simd.md for the encoding, error model, and dispatch strategy.
6
7#[cfg(target_arch = "x86_64")]
8use std::arch::x86_64::*;
9
10#[cfg(target_arch = "aarch64")]
11use std::arch::aarch64::*;
12
13use std::sync::OnceLock;
14
15use super::simd_config;
16
17/// **Unstable**: INT8 quantization parameters; scale/bias scheme may change.
18///
19/// Quantization parameters for int8 conversion.
20#[derive(Debug, Clone, Copy)]
21pub struct QuantizationParams {
22    /// **Unstable**: scale factor; formula may change with scheme update.
23    pub scale: f32,
24    /// **Unstable**: zero point offset; may be removed for symmetric-only quantization.
25    pub zero_point: i8,
26    /// **Unstable**: min float value; may be removed.
27    pub min_val: f32,
28    /// **Unstable**: max float value; may be removed.
29    pub max_val: f32,
30}
31
32impl QuantizationParams {
33    /// **Unstable**: parameter computation; may be folded into `QuantizedVector::from_f32`.
34    ///
35    /// Handles edge cases: empty vectors, NaN, Inf, near-zero vectors.
36    pub fn from_vector(vector: &[f32]) -> Self {
37        // Single pass over finite values to handle NaN/Inf gracefully.
38        let (mut min_val, mut max_val) = minmax_finite(vector);
39
40        // Handle edge case: empty or all non-finite.
41        if !min_val.is_finite() || !max_val.is_finite() {
42            min_val = 0.0;
43            max_val = 0.0;
44        }
45
46        // Symmetric quantization: map [-max_abs, max_abs] to [-127, 127]
47        let max_abs = min_val.abs().max(max_val.abs());
48
49        // Epsilon guard to avoid division by near-zero
50        let scale = if max_abs > 1e-10 {
51            127.0 / max_abs
52        } else {
53            1.0 // All zeros or near-zero case
54        };
55
56        Self {
57            scale,
58            zero_point: 0,
59            min_val,
60            max_val,
61        }
62    }
63}
64
65/// Minimum and maximum over the finite lanes of `v`, in one pass.
66///
67/// Non-finite lanes are skipped. An empty or all-non-finite input yields
68/// `(+inf, -inf)` so the caller's reset branch fires.
69///
70/// Explicit NEON kernel with a scalar fallback rather than a plain guarded
71/// loop: the auto-vectorized form of this reduction was demoted to scalar
72/// code by an unrelated codegen-unit reshuffle (+48% on per-call int8
73/// quantization at 1024 dims), so the hot path must not depend on the
74/// auto-vectorizer's mood.
75fn minmax_finite(v: &[f32]) -> (f32, f32) {
76    #[cfg(target_arch = "x86_64")]
77    {
78        if simd_config().avx2_enabled {
79            // SAFETY: AVX2 was detected at runtime; the kernel bounds every load.
80            return unsafe { minmax_finite_avx2(v) };
81        }
82    }
83    #[cfg(target_arch = "aarch64")]
84    {
85        if simd_config().neon_enabled {
86            // SAFETY: NEON confirmed available at runtime.
87            return unsafe { minmax_finite_neon(v) };
88        }
89    }
90    minmax_finite_scalar(v)
91}
92
93/// Resolves the sign of a zero-valued min/max so every kernel agrees bit-for-bit.
94///
95/// `f32::min`/`f32::max`, NEON's `vminvq_f32`/`vmaxvq_f32`, and AVX2's
96/// `_mm256_min_ps`/`_mm256_max_ps` each return an unspecified one of their operands
97/// when the operands compare equal, so a `-0.0`/`+0.0` tie is broken differently per
98/// kernel and per codegen. Applying IEEE 754-2019 `minimum`/`maximum` ordering
99/// (`-0.0 < +0.0`) to the finished pair makes the choice a stated contract instead of
100/// an artifact: the min takes `-0.0` and the max takes `+0.0` whenever that sign is
101/// present in the input.
102///
103/// The sign scan runs only when a bound is exactly zero, so the cost on a typical
104/// vector is the two comparisons in the guard.
105#[inline]
106fn pin_zero_signs(v: &[f32], min_val: f32, max_val: f32) -> (f32, f32) {
107    if min_val != 0.0 && max_val != 0.0 {
108        return (min_val, max_val);
109    }
110    let mut has_negative_zero = false;
111    let mut has_positive_zero = false;
112    for &value in v {
113        has_negative_zero |= value.to_bits() == (-0.0f32).to_bits();
114        has_positive_zero |= value.to_bits() == 0.0f32.to_bits();
115    }
116    let min_val = if min_val == 0.0 {
117        if has_negative_zero { -0.0 } else { 0.0 }
118    } else {
119        min_val
120    };
121    let max_val = if max_val == 0.0 {
122        if has_positive_zero { 0.0 } else { -0.0 }
123    } else {
124        max_val
125    };
126    (min_val, max_val)
127}
128
129fn minmax_finite_scalar(v: &[f32]) -> (f32, f32) {
130    let mut min_val = f32::INFINITY;
131    let mut max_val = f32::NEG_INFINITY;
132    for &x in v {
133        if x.is_finite() {
134            min_val = min_val.min(x);
135            max_val = max_val.max(x);
136        }
137    }
138    pin_zero_signs(v, min_val, max_val)
139}
140
141#[cfg(test)]
142thread_local! {
143    static I8_MINMAX_SIMD_HITS: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
144}
145
146#[cfg(target_arch = "x86_64")]
147#[target_feature(enable = "avx2")]
148unsafe fn minmax_finite_avx2(v: &[f32]) -> (f32, f32) {
149    #[cfg(test)]
150    I8_MINMAX_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
151
152    let chunks = v.len() / 8;
153    let inf = _mm256_set1_ps(f32::INFINITY);
154    let neg_inf = _mm256_set1_ps(f32::NEG_INFINITY);
155    let sign = _mm256_set1_ps(-0.0);
156    let mut vmin = inf;
157    let mut vmax = neg_inf;
158
159    for i in 0..chunks {
160        let x = _mm256_loadu_ps(v.as_ptr().add(i * 8));
161        let abs = _mm256_andnot_ps(sign, x);
162        let finite = _mm256_cmp_ps(abs, inf, _CMP_LT_OQ);
163        vmin = _mm256_min_ps(vmin, _mm256_blendv_ps(inf, x, finite));
164        vmax = _mm256_max_ps(vmax, _mm256_blendv_ps(neg_inf, x, finite));
165    }
166
167    let mut min_lanes = [0.0f32; 8];
168    let mut max_lanes = [0.0f32; 8];
169    _mm256_storeu_ps(min_lanes.as_mut_ptr(), vmin);
170    _mm256_storeu_ps(max_lanes.as_mut_ptr(), vmax);
171
172    let mut min_val = min_lanes.into_iter().fold(f32::INFINITY, f32::min);
173    let mut max_val = max_lanes.into_iter().fold(f32::NEG_INFINITY, f32::max);
174    for &x in &v[chunks * 8..] {
175        if x.is_finite() {
176            min_val = min_val.min(x);
177            max_val = max_val.max(x);
178        }
179    }
180    pin_zero_signs(v, min_val, max_val)
181}
182
183#[cfg(target_arch = "aarch64")]
184unsafe fn minmax_finite_neon(v: &[f32]) -> (f32, f32) {
185    #[cfg(test)]
186    I8_MINMAX_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
187
188    let chunks = v.len() / 4;
189    let inf = unsafe { vdupq_n_f32(f32::INFINITY) };
190    let neg_inf = unsafe { vdupq_n_f32(f32::NEG_INFINITY) };
191    let mut vmin = inf;
192    let mut vmax = neg_inf;
193    for i in 0..chunks {
194        // SAFETY: `i * 4 + 3 < v.len()` by the chunk bound.
195        let x = unsafe { vld1q_f32(v.as_ptr().add(i * 4)) };
196        unsafe {
197            // `|x| < +inf` holds exactly for finite lanes (false for NaN, ±inf),
198            // so masked lanes contribute identity elements, same as skipping.
199            let finite = vcaltq_f32(x, inf);
200            vmin = vminq_f32(vmin, vbslq_f32(finite, x, inf));
201            vmax = vmaxq_f32(vmax, vbslq_f32(finite, x, neg_inf));
202        }
203    }
204    let (mut min_val, mut max_val) = unsafe { (vminvq_f32(vmin), vmaxvq_f32(vmax)) };
205    for &x in &v[chunks * 4..] {
206        if x.is_finite() {
207            min_val = min_val.min(x);
208            max_val = max_val.max(x);
209        }
210    }
211    pin_zero_signs(v, min_val, max_val)
212}
213
214/// **Unstable**: INT8 quantized vector; struct layout and invariants may change.
215///
216/// Quantized int8 vector with its parameters.
217#[derive(Debug, Clone)]
218pub struct QuantizedVector {
219    /// Invariant: all values in `[-127, 127]`. Enforced by `from_f32` clamping.
220    /// Private — the invariant makes release-mode assert scans unnecessary.
221    data: Vec<i8>,
222    /// **Unstable**: quantization parameters; may be separated from the vector.
223    pub params: QuantizationParams,
224    /// **Unstable**: L2 norm; may be removed or moved.
225    pub norm: f32,
226}
227
228impl QuantizedVector {
229    /// Returns the quantized data as a slice. All values are in `[-127, 127]`.
230    #[inline]
231    pub fn data(&self) -> &[i8] {
232        &self.data
233    }
234
235    /// Returns the number of quantized elements.
236    #[inline]
237    pub fn len(&self) -> usize {
238        self.data.len()
239    }
240
241    /// Returns `true` if the quantized vector has no elements.
242    #[inline]
243    pub fn is_empty(&self) -> bool {
244        self.data.is_empty()
245    }
246}
247
248impl QuantizedVector {
249    /// **Unstable**: quantization constructor; clamping behavior may change.
250    pub fn from_f32(vector: &[f32]) -> Self {
251        let mut params = QuantizationParams::from_vector(vector);
252
253        // Defensive guard: avoid NaN/Inf/zero scale.
254        if !params.scale.is_finite() || params.scale == 0.0 {
255            params.scale = 1.0;
256        }
257
258        // Compute L2 norm of finite values (NaN/Inf are treated as 0.0).
259        let mut norm_sq = 0.0f32;
260        for &v in vector {
261            if v.is_finite() {
262                norm_sq += v * v;
263            }
264        }
265        let norm = norm_sq.sqrt();
266
267        let data = quantize_i8(vector, params.scale);
268
269        Self { data, params, norm }
270    }
271
272    /// **Unstable**: dequantizes this vector using its stored scale.
273    ///
274    /// See [`docs/simd.md`](../../docs/simd.md#int8-vectors) for the encoding and error bounds.
275    pub fn to_f32(&self) -> Vec<f32> {
276        let scale = if self.params.scale.is_finite() && self.params.scale != 0.0 {
277            self.params.scale
278        } else {
279            1.0
280        };
281
282        self.data.iter().map(|&v| v as f32 / scale).collect()
283    }
284
285    /// **Unstable**: delegates to `dot_product_i8`; SIMD dispatch may change.
286    #[inline]
287    pub fn dot_product(&self, other: &QuantizedVector) -> f32 {
288        dot_product_i8(self, other)
289    }
290
291    /// **Unstable**: delegates to `cosine_similarity_i8`; SIMD dispatch may change.
292    #[inline]
293    pub fn cosine_similarity(&self, other: &QuantizedVector) -> f32 {
294        cosine_similarity_i8(self, other)
295    }
296}
297
298fn quantize_i8(vector: &[f32], scale: f32) -> Vec<i8> {
299    #[cfg(target_arch = "x86_64")]
300    {
301        if simd_config().avx2_enabled {
302            // SAFETY: AVX2 was detected at runtime; the kernel bounds every load and store.
303            return unsafe { quantize_i8_avx2(vector, scale) };
304        }
305    }
306    #[cfg(target_arch = "aarch64")]
307    {
308        if simd_config().neon_enabled {
309            // SAFETY: NEON was detected at runtime; the kernel bounds every load and store.
310            return unsafe { quantize_i8_neon(vector, scale) };
311        }
312    }
313    quantize_i8_scalar(vector, scale)
314}
315
316fn quantize_i8_scalar(vector: &[f32], scale: f32) -> Vec<i8> {
317    vector
318        .iter()
319        .map(|&v| quantize_i8_value(v, scale))
320        .collect()
321}
322
323#[inline]
324fn quantize_i8_value(value: f32, scale: f32) -> i8 {
325    if value.is_finite() {
326        (value * scale).round().clamp(-127.0, 127.0) as i8
327    } else {
328        0
329    }
330}
331
332#[cfg(test)]
333thread_local! {
334    static I8_QUANTIZE_SIMD_HITS: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
335}
336
337#[cfg(target_arch = "x86_64")]
338#[target_feature(enable = "avx2")]
339unsafe fn quantize_i8_avx2(vector: &[f32], scale: f32) -> Vec<i8> {
340    #[cfg(test)]
341    I8_QUANTIZE_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
342
343    let mut data = vec![0i8; vector.len()];
344    let chunks = vector.len() / 8;
345    let scale_scalar = scale;
346    let scale = _mm256_set1_ps(scale_scalar);
347    let inf = _mm256_set1_ps(f32::INFINITY);
348    let sign = _mm256_set1_ps(-0.0);
349    let low = _mm256_set1_ps(-127.0);
350    let high = _mm256_set1_ps(127.0);
351    let half = _mm256_set1_ps(0.5);
352    let negative_half = _mm256_set1_ps(-0.5);
353    let one = _mm256_set1_epi32(1);
354    let negative_one = _mm256_set1_epi32(-1);
355
356    for i in 0..chunks {
357        let base = i * 8;
358        let input = _mm256_loadu_ps(vector.as_ptr().add(base));
359        let abs = _mm256_andnot_ps(sign, input);
360        let finite = _mm256_cmp_ps(abs, inf, _CMP_LT_OQ);
361        let values = _mm256_and_ps(input, finite);
362        let scaled = _mm256_mul_ps(values, scale);
363        let clamped = _mm256_min_ps(_mm256_max_ps(scaled, low), high);
364        let truncated = _mm256_cvttps_epi32(clamped);
365        let fraction = _mm256_sub_ps(clamped, _mm256_cvtepi32_ps(truncated));
366        let round_up = _mm256_castps_si256(_mm256_cmp_ps(fraction, half, _CMP_GE_OQ));
367        let round_down = _mm256_castps_si256(_mm256_cmp_ps(fraction, negative_half, _CMP_LE_OQ));
368        let rounded = _mm256_add_epi32(
369            _mm256_add_epi32(truncated, _mm256_and_si256(round_up, one)),
370            _mm256_and_si256(round_down, negative_one),
371        );
372        let mut lanes = [0i32; 8];
373        _mm256_storeu_si256(lanes.as_mut_ptr().cast::<__m256i>(), rounded);
374        for (offset, lane) in lanes.into_iter().enumerate() {
375            data[base + offset] = lane as i8;
376        }
377    }
378
379    for i in chunks * 8..vector.len() {
380        data[i] = quantize_i8_value(vector[i], scale_scalar);
381    }
382    data
383}
384
385#[cfg(target_arch = "aarch64")]
386#[target_feature(enable = "neon")]
387unsafe fn quantize_i8_neon(vector: &[f32], scale: f32) -> Vec<i8> {
388    #[cfg(test)]
389    I8_QUANTIZE_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
390
391    let mut data = vec![0i8; vector.len()];
392    let chunks = vector.len() / 4;
393    let scale_vector = vdupq_n_f32(scale);
394    let inf = vdupq_n_f32(f32::INFINITY);
395    let zero = vdupq_n_f32(0.0);
396    let low = vdupq_n_f32(-127.0);
397    let high = vdupq_n_f32(127.0);
398
399    for i in 0..chunks {
400        let base = i * 4;
401        let input = vld1q_f32(vector.as_ptr().add(base));
402        let finite = vcaltq_f32(input, inf);
403        let values = vbslq_f32(finite, input, zero);
404        let scaled = vmulq_f32(values, scale_vector);
405        let clamped = vminq_f32(vmaxq_f32(scaled, low), high);
406        let rounded = vcvtaq_s32_f32(clamped);
407        let mut lanes = [0i32; 4];
408        vst1q_s32(lanes.as_mut_ptr(), rounded);
409        for (offset, lane) in lanes.into_iter().enumerate() {
410            data[base + offset] = lane as i8;
411        }
412    }
413
414    for i in chunks * 4..vector.len() {
415        data[i] = quantize_i8_value(vector[i], scale);
416    }
417    data
418}
419
420/// **Unstable**: computes an approximate float dot product, returning `0.0` for a mismatch.
421///
422/// Inputs satisfy the constructor-owned `[-127, 127]` invariant.
423/// See [`docs/simd.md`](../../docs/simd.md#raw-int8-input-invariant) for its SIMD requirement.
424#[inline]
425pub fn dot_product_i8(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
426    debug_assert!(a.data.iter().all(|&v| v != -128i8));
427    debug_assert!(b.data.iter().all(|&v| v != -128i8));
428
429    if a.data.len() != b.data.len() {
430        return 0.0;
431    }
432
433    let denom = a.params.scale * b.params.scale;
434    if denom == 0.0 || !denom.is_finite() {
435        return 0.0;
436    }
437
438    dot_product_i8_dispatch(&a.data, &b.data) / denom
439}
440
441/// Trusted INT8 dot product for constructor-owned vectors in prepared-query paths.
442///
443/// Uses `debug_assert!` instead of `assert!`; callers must guarantee vectors
444/// were produced by `QuantizedVector::from_f32` or equivalent (clamped to [-127,127]).
445#[inline]
446pub(crate) fn dot_product_i8_trusted(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
447    if a.data.len() != b.data.len() {
448        return 0.0;
449    }
450    let denom = a.params.scale * b.params.scale;
451    if denom == 0.0 || !denom.is_finite() {
452        return 0.0;
453    }
454    debug_assert!(a.data.iter().all(|&v| v != i8::MIN));
455    debug_assert!(b.data.iter().all(|&v| v != i8::MIN));
456    dot_product_i8_dispatch(&a.data, &b.data) / denom
457}
458
459/// **Unstable**: SIMD INT8 cosine similarity; norm storage approach may change.
460///
461/// Uses pre-computed norms for efficiency.
462#[inline]
463pub fn cosine_similarity_i8(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
464    let denom = a.norm * b.norm;
465    if denom == 0.0 || !denom.is_finite() {
466        return 0.0;
467    }
468    dot_product_i8(a, b) / denom
469}
470
471/// Computes INT8 cosine similarity for constructor-owned vectors without a release scan.
472///
473/// See [`docs/simd.md`](../../docs/simd.md#raw-int8-input-invariant) for the trusted-path precondition.
474#[inline]
475pub(crate) fn cosine_similarity_i8_trusted(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
476    let denom = a.norm * b.norm;
477    if denom == 0.0 || !denom.is_finite() {
478        return 0.0;
479    }
480    dot_product_i8_trusted(a, b) / denom
481}
482
483/// Computes an INT8 dot product with FEAT_DotProd and guarded prefetch.
484///
485/// # Safety
486/// Caller must provide FEAT_DotProd, equal `[-127, 127]` slices, and bounded prefetches.
487/// See [`docs/simd.md`](../../docs/simd.md#int8-vectors) for dispatch and implementation details.
488#[cfg(target_arch = "aarch64")]
489#[target_feature(enable = "dotprod")]
490unsafe fn dot_product_i8_neon_unrolled(a: &[i8], b: &[i8]) -> f32 {
491    const SIMD_WIDTH: usize = 16;
492    const UNROLL: usize = 4;
493    const CHUNK_SIZE: usize = SIMD_WIDTH * UNROLL;
494    const PREFETCH_DISTANCE: usize = CHUNK_SIZE;
495    let n = a.len();
496    debug_assert_eq!(n, b.len());
497    let chunks = n / CHUNK_SIZE;
498
499    let mut sum0 = vdupq_n_s32(0);
500    let mut sum1 = vdupq_n_s32(0);
501    let mut sum2 = vdupq_n_s32(0);
502    let mut sum3 = vdupq_n_s32(0);
503
504    // SDOT is selected only after runtime FEAT_DotProd detection — see docs/simd.md.
505    for i in 0..chunks {
506        let base = i * CHUNK_SIZE;
507
508        let next_base = base + PREFETCH_DISTANCE;
509        if next_base + CHUNK_SIZE <= n {
510            core::arch::asm!(
511                "prfm pldl1keep, [{ptr}]",
512                ptr = in(reg) a.as_ptr().add(next_base),
513                options(nostack, readonly, preserves_flags)
514            );
515            core::arch::asm!(
516                "prfm pldl1keep, [{ptr}]",
517                ptr = in(reg) b.as_ptr().add(next_base),
518                options(nostack, readonly, preserves_flags)
519            );
520        }
521
522        let a0 = vld1q_s8(a.as_ptr().add(base));
523        let b0 = vld1q_s8(b.as_ptr().add(base));
524        let a1 = vld1q_s8(a.as_ptr().add(base + SIMD_WIDTH));
525        let b1 = vld1q_s8(b.as_ptr().add(base + SIMD_WIDTH));
526        let a2 = vld1q_s8(a.as_ptr().add(base + SIMD_WIDTH * 2));
527        let b2 = vld1q_s8(b.as_ptr().add(base + SIMD_WIDTH * 2));
528        let a3 = vld1q_s8(a.as_ptr().add(base + SIMD_WIDTH * 3));
529        let b3 = vld1q_s8(b.as_ptr().add(base + SIMD_WIDTH * 3));
530
531        core::arch::asm!(
532            "sdot {s0:v}.4s, {a0:v}.16b, {b0:v}.16b",
533            "sdot {s1:v}.4s, {a1:v}.16b, {b1:v}.16b",
534            "sdot {s2:v}.4s, {a2:v}.16b, {b2:v}.16b",
535            "sdot {s3:v}.4s, {a3:v}.16b, {b3:v}.16b",
536            s0 = inout(vreg) sum0,
537            a0 = in(vreg) a0,
538            b0 = in(vreg) b0,
539            s1 = inout(vreg) sum1,
540            a1 = in(vreg) a1,
541            b1 = in(vreg) b1,
542            s2 = inout(vreg) sum2,
543            a2 = in(vreg) a2,
544            b2 = in(vreg) b2,
545            s3 = inout(vreg) sum3,
546            a3 = in(vreg) a3,
547            b3 = in(vreg) b3,
548            options(nomem, nostack, preserves_flags)
549        );
550    }
551
552    let sum01 = vaddq_s32(sum0, sum1);
553    let sum23 = vaddq_s32(sum2, sum3);
554    let mut sum_vec = vaddq_s32(sum01, sum23);
555
556    // Tail: remaining full 16-byte vectors using sdot
557    let tail_start = chunks * CHUNK_SIZE;
558    let tail_chunks = (n - tail_start) / SIMD_WIDTH;
559    for j in 0..tail_chunks {
560        let base = tail_start + j * SIMD_WIDTH;
561        let at = vld1q_s8(a.as_ptr().add(base));
562        let bt = vld1q_s8(b.as_ptr().add(base));
563        core::arch::asm!(
564            "sdot {acc:v}.4s, {a:v}.16b, {b:v}.16b",
565            acc = inout(vreg) sum_vec,
566            a = in(vreg) at,
567            b = in(vreg) bt,
568            options(nomem, nostack, preserves_flags)
569        );
570    }
571
572    let sum = vaddvq_s32(sum_vec);
573
574    // Scalar tail: only the final < SIMD_WIDTH elements
575    let remainder_start = tail_start + tail_chunks * SIMD_WIDTH;
576    let remainder: i32 = a[remainder_start..]
577        .iter()
578        .zip(b[remainder_start..].iter())
579        .map(|(&x, &y)| x as i32 * y as i32)
580        .sum();
581
582    (sum + remainder) as f32
583}
584
585/// Emulate `mm512_sign_epi8(b, a)` which doesn't exist in AVX-512.
586///
587/// Returns: b[i] if a[i] > 0, -b[i] if a[i] < 0, 0 if a[i] == 0.
588///
589/// # Safety
590/// Requires AVX-512BW.
591#[cfg(target_arch = "x86_64")]
592#[target_feature(enable = "avx512f", enable = "avx512bw")]
593#[inline]
594unsafe fn mm512_sign_epi8(b: __m512i, a: __m512i) -> __m512i {
595    let zero = _mm512_setzero_si512();
596    let neg_b = _mm512_sub_epi8(zero, b);
597    // mask where a < 0
598    let mask_neg = _mm512_cmplt_epi8_mask(a, zero);
599    // mask where a == 0
600    let mask_zero = _mm512_cmpeq_epi8_mask(a, zero);
601    // Start with b, replace with -b where a < 0
602    let result = _mm512_mask_blend_epi8(mask_neg, b, neg_b);
603    // Replace with 0 where a == 0
604    _mm512_mask_blend_epi8(mask_zero, result, zero)
605}
606
607/// Computes an INT8 dot product with AVX-512 VNNI.
608///
609/// # Safety
610/// Caller must provide AVX-512F/VNNI/BW and equal `[-127, 127]` slices; bounds are chunked.
611/// See [`docs/simd.md`](../../docs/simd.md#int8-vectors) for signed-product transformation details.
612#[cfg(target_arch = "x86_64")]
613#[target_feature(enable = "avx512f", enable = "avx512vnni", enable = "avx512bw")]
614unsafe fn dot_product_i8_avx512vnni(a: &[i8], b: &[i8]) -> f32 {
615    const SIMD_WIDTH: usize = 64; // 64 int8s per 512-bit register
616    const UNROLL: usize = 4;
617    const CHUNK_SIZE: usize = SIMD_WIDTH * UNROLL;
618    let n = a.len();
619    debug_assert_eq!(n, b.len());
620    debug_assert!(a.iter().all(|&v| v != i8::MIN));
621    debug_assert!(b.iter().all(|&v| v != i8::MIN));
622    let chunks = n / CHUNK_SIZE;
623
624    // 4 independent int32 accumulators (16 int32s each)
625    let mut sum0 = _mm512_setzero_si512();
626    let mut sum1 = _mm512_setzero_si512();
627    let mut sum2 = _mm512_setzero_si512();
628    let mut sum3 = _mm512_setzero_si512();
629
630    for i in 0..chunks {
631        let base = i * CHUNK_SIZE;
632
633        // VNNI: dpbusd computes sum += a[unsigned] * b[signed]
634        // For signed * signed, we use: abs(a) * sign(b, a)
635        let a0 = _mm512_loadu_si512(a.as_ptr().add(base) as *const __m512i);
636        let b0 = _mm512_loadu_si512(b.as_ptr().add(base) as *const __m512i);
637        let a0_abs = _mm512_abs_epi8(a0);
638        let b0_signed = mm512_sign_epi8(b0, a0);
639        sum0 = _mm512_dpbusd_epi32(sum0, a0_abs, b0_signed);
640
641        let a1 = _mm512_loadu_si512(a.as_ptr().add(base + SIMD_WIDTH) as *const __m512i);
642        let b1 = _mm512_loadu_si512(b.as_ptr().add(base + SIMD_WIDTH) as *const __m512i);
643        let a1_abs = _mm512_abs_epi8(a1);
644        let b1_signed = mm512_sign_epi8(b1, a1);
645        sum1 = _mm512_dpbusd_epi32(sum1, a1_abs, b1_signed);
646
647        let a2 = _mm512_loadu_si512(a.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m512i);
648        let b2 = _mm512_loadu_si512(b.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m512i);
649        let a2_abs = _mm512_abs_epi8(a2);
650        let b2_signed = mm512_sign_epi8(b2, a2);
651        sum2 = _mm512_dpbusd_epi32(sum2, a2_abs, b2_signed);
652
653        let a3 = _mm512_loadu_si512(a.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m512i);
654        let b3 = _mm512_loadu_si512(b.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m512i);
655        let a3_abs = _mm512_abs_epi8(a3);
656        let b3_signed = mm512_sign_epi8(b3, a3);
657        sum3 = _mm512_dpbusd_epi32(sum3, a3_abs, b3_signed);
658    }
659
660    // Combine accumulators
661    let sum01 = _mm512_add_epi32(sum0, sum1);
662    let sum23 = _mm512_add_epi32(sum2, sum3);
663    let sum_vec = _mm512_add_epi32(sum01, sum23);
664
665    // Horizontal sum of 16 int32s
666    let sum = _mm512_reduce_add_epi32(sum_vec);
667
668    // Handle remainder with scalar
669    let remainder_start = chunks * CHUNK_SIZE;
670    let remainder: i32 = a[remainder_start..]
671        .iter()
672        .zip(b[remainder_start..].iter())
673        .map(|(&x, &y)| x as i32 * y as i32)
674        .sum();
675
676    (sum + remainder) as f32
677}
678
679/// Computes an INT8 dot product with AVX2 and guarded prefetch.
680///
681/// # Safety
682/// Caller must provide AVX2 and equal `[-127, 127]` slices; bounds and prefetch are guarded.
683/// See [`docs/simd.md`](../../docs/simd.md#int8-vectors) for signed-product transformation details.
684#[cfg(target_arch = "x86_64")]
685#[target_feature(enable = "avx2")]
686unsafe fn dot_product_i8_avx2_unrolled(a: &[i8], b: &[i8]) -> f32 {
687    const SIMD_WIDTH: usize = 32;
688    const UNROLL: usize = 4;
689    const CHUNK_SIZE: usize = SIMD_WIDTH * UNROLL;
690    // Prefetch one full chunk ahead for both input arrays.
691    const PREFETCH_DISTANCE: usize = CHUNK_SIZE;
692    let n = a.len();
693    debug_assert_eq!(n, b.len());
694    debug_assert!(a.iter().all(|&v| v != i8::MIN));
695    debug_assert!(b.iter().all(|&v| v != i8::MIN));
696    let chunks = n / CHUNK_SIZE;
697
698    // 4 independent int32 accumulators
699    let mut sum0 = _mm256_setzero_si256();
700    let mut sum1 = _mm256_setzero_si256();
701    let mut sum2 = _mm256_setzero_si256();
702    let mut sum3 = _mm256_setzero_si256();
703
704    let ones = _mm256_set1_epi16(1);
705
706    for i in 0..chunks {
707        let base = i * CHUNK_SIZE;
708
709        // Software prefetch for the next chunk.
710        let next_base = base + PREFETCH_DISTANCE;
711        if next_base + CHUNK_SIZE <= n {
712            _mm_prefetch(a.as_ptr().add(next_base), _MM_HINT_T0);
713            _mm_prefetch(b.as_ptr().add(next_base), _MM_HINT_T0);
714        }
715
716        // Unroll 0
717        let a0 = _mm256_loadu_si256(a.as_ptr().add(base) as *const __m256i);
718        let b0 = _mm256_loadu_si256(b.as_ptr().add(base) as *const __m256i);
719        let prod0 = _mm256_maddubs_epi16(_mm256_abs_epi8(a0), _mm256_sign_epi8(b0, a0));
720        let prod0_32 = _mm256_madd_epi16(prod0, ones);
721        sum0 = _mm256_add_epi32(sum0, prod0_32);
722
723        // Unroll 1
724        let a1 = _mm256_loadu_si256(a.as_ptr().add(base + SIMD_WIDTH) as *const __m256i);
725        let b1 = _mm256_loadu_si256(b.as_ptr().add(base + SIMD_WIDTH) as *const __m256i);
726        let prod1 = _mm256_maddubs_epi16(_mm256_abs_epi8(a1), _mm256_sign_epi8(b1, a1));
727        let prod1_32 = _mm256_madd_epi16(prod1, ones);
728        sum1 = _mm256_add_epi32(sum1, prod1_32);
729
730        // Unroll 2
731        let a2 = _mm256_loadu_si256(a.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m256i);
732        let b2 = _mm256_loadu_si256(b.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m256i);
733        let prod2 = _mm256_maddubs_epi16(_mm256_abs_epi8(a2), _mm256_sign_epi8(b2, a2));
734        let prod2_32 = _mm256_madd_epi16(prod2, ones);
735        sum2 = _mm256_add_epi32(sum2, prod2_32);
736
737        // Unroll 3
738        let a3 = _mm256_loadu_si256(a.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m256i);
739        let b3 = _mm256_loadu_si256(b.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m256i);
740        let prod3 = _mm256_maddubs_epi16(_mm256_abs_epi8(a3), _mm256_sign_epi8(b3, a3));
741        let prod3_32 = _mm256_madd_epi16(prod3, ones);
742        sum3 = _mm256_add_epi32(sum3, prod3_32);
743    }
744
745    // Combine accumulators
746    let sum01 = _mm256_add_epi32(sum0, sum1);
747    let sum23 = _mm256_add_epi32(sum2, sum3);
748    let sum_vec = _mm256_add_epi32(sum01, sum23);
749
750    // Horizontal sum
751    let sum128_lo = _mm256_castsi256_si128(sum_vec);
752    let sum128_hi = _mm256_extracti128_si256(sum_vec, 1);
753    let sum128 = _mm_add_epi32(sum128_lo, sum128_hi);
754    let sum64 = _mm_add_epi32(sum128, _mm_srli_si128(sum128, 8));
755    let sum32 = _mm_add_epi32(sum64, _mm_srli_si128(sum64, 4));
756    let sum = _mm_cvtsi128_si32(sum32);
757
758    // Handle remainder
759    let remainder_start = chunks * CHUNK_SIZE;
760    let remainder: i32 = a[remainder_start..]
761        .iter()
762        .zip(b[remainder_start..].iter())
763        .map(|(&x, &y)| x as i32 * y as i32)
764        .sum();
765
766    (sum + remainder) as f32
767}
768
769// ============================================================================
770// INT8 kernel dispatch cache (mirrors f32 DotKernel pattern in dot_product.rs)
771// ============================================================================
772
773/// INT8 dot-product kernel function pointer type.
774pub type I8DotKernel = fn(&[i8], &[i8]) -> f32;
775
776static I8_DOT_KERNEL: OnceLock<I8DotKernel> = OnceLock::new();
777
778/// Return the cached INT8 dot-product kernel for tight loops.
779#[inline]
780pub fn resolved_i8_dot_kernel() -> I8DotKernel {
781    *I8_DOT_KERNEL.get_or_init(resolve_i8_dot_kernel)
782}
783
784fn resolve_i8_dot_kernel() -> I8DotKernel {
785    let config = simd_config();
786
787    #[cfg(target_arch = "aarch64")]
788    {
789        // The NEON kernel uses SDOT (ARMv8.2 FEAT_DotProd), which is optional on
790        // Armv8.2/v8.3. Only dispatch when dotprod_enabled is confirmed at runtime.
791        if config.neon_enabled && config.dotprod_enabled {
792            return dot_product_i8_neon_kernel;
793        }
794    }
795
796    #[cfg(target_arch = "x86_64")]
797    {
798        if config.avx512vnni_enabled {
799            return dot_product_i8_avx512vnni_kernel;
800        }
801        if config.avx2_enabled {
802            return dot_product_i8_avx2_kernel;
803        }
804    }
805
806    dot_product_i8_scalar_kernel
807}
808
809#[cfg(target_arch = "aarch64")]
810fn dot_product_i8_neon_kernel(a: &[i8], b: &[i8]) -> f32 {
811    // SAFETY: stored only when NEON+dotprod detected at init time.
812    unsafe { dot_product_i8_neon_unrolled(a, b) }
813}
814
815#[cfg(target_arch = "x86_64")]
816fn dot_product_i8_avx512vnni_kernel(a: &[i8], b: &[i8]) -> f32 {
817    debug_assert!(a.iter().all(|&v| v != i8::MIN));
818    debug_assert!(b.iter().all(|&v| v != i8::MIN));
819    // SAFETY: stored only when AVX-512F+VNNI+BW were detected at init time.
820    unsafe { dot_product_i8_avx512vnni(a, b) }
821}
822
823#[cfg(target_arch = "x86_64")]
824fn dot_product_i8_avx2_kernel(a: &[i8], b: &[i8]) -> f32 {
825    debug_assert!(a.iter().all(|&v| v != i8::MIN));
826    debug_assert!(b.iter().all(|&v| v != i8::MIN));
827    // SAFETY: stored only when AVX2 was detected at init time.
828    unsafe { dot_product_i8_avx2_unrolled(a, b) }
829}
830
831fn dot_product_i8_scalar_kernel(a: &[i8], b: &[i8]) -> f32 {
832    a.iter()
833        .zip(b.iter())
834        .map(|(&x, &y)| x as i32 * y as i32)
835        .sum::<i32>() as f32
836}
837
838/// Dispatch a validated raw INT8 dot product.
839#[inline]
840fn dot_product_i8_dispatch(a: &[i8], b: &[i8]) -> f32 {
841    resolved_i8_dot_kernel()(a, b)
842}
843
844/// **Unstable**: computes an unscaled raw INT8 dot product, returning `0.0` for a mismatch.
845///
846/// Every value must lie in `[-127, 127]`; `i8::MIN` is numerically invalid.
847/// See [`docs/simd.md`](../../docs/simd.md#raw-int8-input-invariant) for the release-mode precondition.
848#[inline]
849pub fn dot_product_i8_raw(a: &[i8], b: &[i8]) -> f32 {
850    if a.len() != b.len() {
851        return 0.0;
852    }
853    debug_assert!(
854        a.iter().all(|&v| v != -128i8),
855        "dot_product_i8_raw: slice a contains -128, violating the [-127, 127] SIMD invariant"
856    );
857    debug_assert!(
858        b.iter().all(|&v| v != -128i8),
859        "dot_product_i8_raw: slice b contains -128, violating the [-127, 127] SIMD invariant"
860    );
861    dot_product_i8_dispatch(a, b)
862}
863
864#[cfg(test)]
865mod simd_parity_tests {
866    use super::*;
867
868    fn gen_vec(dim: usize, seed: u64) -> Vec<f32> {
869        let mut state = seed ^ ((dim as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
870        (0..dim)
871            .map(|i| {
872                state = state
873                    .wrapping_mul(6364136223846793005)
874                    .wrapping_add(1442695040888963407)
875                    .wrapping_add(i as u64);
876                let unit = ((state >> 32) as u32) as f32 / u32::MAX as f32;
877                unit * 2.0 - 1.0
878            })
879            .collect()
880    }
881
882    #[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
883    #[test]
884    fn test_i8_quantize_explicit_simd_matches_scalar_and_is_dispatched() {
885        #[cfg(target_arch = "x86_64")]
886        if !std::arch::is_x86_feature_detected!("avx2") {
887            return;
888        }
889
890        for dim in [0usize, 1, 3, 4, 7, 8, 9, 31, 32, 33, 383, 384, 385] {
891            let mut input = gen_vec(dim, 900 + dim as u64);
892            if dim > 0 {
893                input[0] = f32::NAN;
894            }
895            if dim > 1 {
896                input[1] = f32::INFINITY;
897            }
898            if dim > 2 {
899                input[2] = f32::NEG_INFINITY;
900            }
901            if dim > 3 {
902                input[3] = 0.25;
903            }
904            if dim > 4 {
905                input[4] = -0.25;
906            }
907            if dim > 5 {
908                input[5] = f32::from_bits(0.25f32.to_bits() - 1);
909            }
910            if dim > 6 {
911                input[6] = f32::from_bits(0.25f32.to_bits() + 1);
912            }
913
914            let scalar = quantize_i8_scalar(&input, 2.0);
915            #[cfg(target_arch = "aarch64")]
916            // SAFETY: baseline aarch64 provides NEON; the kernel bounds every access.
917            let simd = unsafe { quantize_i8_neon(&input, 2.0) };
918            #[cfg(target_arch = "x86_64")]
919            // SAFETY: AVX2 was detected above; the kernel bounds every access.
920            let simd = unsafe { quantize_i8_avx2(&input, 2.0) };
921            assert_eq!(simd, scalar, "explicit SIMD mismatch at dim={dim}");
922        }
923
924        let input = gen_vec(385, 1_063);
925        let before = I8_QUANTIZE_SIMD_HITS.with(std::cell::Cell::get);
926        let quantized = QuantizedVector::from_f32(&input);
927        let after = I8_QUANTIZE_SIMD_HITS.with(std::cell::Cell::get);
928        assert_eq!(
929            after,
930            before + 1,
931            "QuantizedVector::from_f32 did not execute its explicit SIMD quantizer"
932        );
933        assert_eq!(
934            quantized.data,
935            quantize_i8_scalar(&input, quantized.params.scale)
936        );
937    }
938
939    #[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
940    #[test]
941    fn test_finite_minmax_explicit_simd_matches_scalar() {
942        #[cfg(target_arch = "x86_64")]
943        if !std::arch::is_x86_feature_detected!("avx2") {
944            return;
945        }
946
947        for dim in [0usize, 1, 3, 4, 7, 8, 9, 31, 32, 33, 383, 384, 385] {
948            let mut input = gen_vec(dim, 1_100 + dim as u64);
949            if dim > 0 {
950                input[0] = f32::NAN;
951            }
952            if dim > 1 {
953                input[1] = f32::INFINITY;
954            }
955            if dim > 2 {
956                input[2] = f32::NEG_INFINITY;
957            }
958            if dim > 3 {
959                input[3] = -0.0;
960            }
961            if dim > 4 {
962                input[4] = 0.0;
963            }
964
965            let scalar = minmax_finite_scalar(&input);
966            #[cfg(target_arch = "aarch64")]
967            // SAFETY: baseline aarch64 provides NEON; the kernel bounds every access.
968            let simd = unsafe { minmax_finite_neon(&input) };
969            #[cfg(target_arch = "x86_64")]
970            // SAFETY: AVX2 was detected above; the kernel bounds every access.
971            let simd = unsafe { minmax_finite_avx2(&input) };
972            assert_eq!(
973                (simd.0.to_bits(), simd.1.to_bits()),
974                (scalar.0.to_bits(), scalar.1.to_bits()),
975                "finite min/max mismatch at dim={dim}"
976            );
977        }
978
979        let mut min_zero_input = vec![1.0f32; 16];
980        min_zero_input[0] = -0.0;
981        min_zero_input[8] = 0.0;
982        let mut max_zero_input = vec![-1.0f32; 16];
983        max_zero_input[0] = 0.0;
984        max_zero_input[8] = -0.0;
985        for input in [&min_zero_input, &max_zero_input] {
986            let scalar = minmax_finite_scalar(input);
987            #[cfg(target_arch = "aarch64")]
988            // SAFETY: baseline aarch64 provides NEON; the kernel bounds every access.
989            let simd = unsafe { minmax_finite_neon(input) };
990            #[cfg(target_arch = "x86_64")]
991            // SAFETY: AVX2 was detected above; the kernel bounds every access.
992            let simd = unsafe { minmax_finite_avx2(input) };
993            assert_eq!(
994                (simd.0.to_bits(), simd.1.to_bits()),
995                (scalar.0.to_bits(), scalar.1.to_bits()),
996                "finite min/max signed-zero mismatch"
997            );
998        }
999
1000        let input = gen_vec(385, 1_063);
1001        let before = I8_MINMAX_SIMD_HITS.with(std::cell::Cell::get);
1002        let params = QuantizationParams::from_vector(&input);
1003        let after = I8_MINMAX_SIMD_HITS.with(std::cell::Cell::get);
1004        assert_eq!(
1005            after,
1006            before + 1,
1007            "QuantizationParams::from_vector did not execute its explicit SIMD reducer"
1008        );
1009        let scalar = minmax_finite_scalar(&input);
1010        assert_eq!(
1011            (params.min_val.to_bits(), params.max_val.to_bits()),
1012            (scalar.0.to_bits(), scalar.1.to_bits())
1013        );
1014    }
1015
1016    /// The signed-zero tie-break is a contract, not whatever the reduction happened to
1017    /// pick. Parity tests can only compare kernels against each other, so they pass on
1018    /// any architecture whose kernels agree by accident; this pins the value itself and
1019    /// runs everywhere, including targets with no explicit SIMD path at all.
1020    ///
1021    /// The first block drives `pin_zero_signs` directly with bounds carrying the wrong
1022    /// sign. That matters: on aarch64 the scalar fold already returns the contracted
1023    /// signs, so assertions routed through `minmax_finite_scalar` alone still pass with
1024    /// the sign pass deleted. Only x86-64 breaks the tie the other way, and a guard that
1025    /// can be removed without any local test noticing is not a guard.
1026    #[test]
1027    fn test_minmax_finite_pins_zero_signs() {
1028        let both_signs = [-0.0f32, 0.0, -1.0];
1029        assert_eq!(
1030            pin_zero_signs(&both_signs, -1.0, -0.0).1.to_bits(),
1031            0.0f32.to_bits(),
1032            "a max of -0.0 must be rewritten to +0.0 when +0.0 is present"
1033        );
1034        assert_eq!(
1035            pin_zero_signs(&both_signs, 0.0, -1.0).0.to_bits(),
1036            (-0.0f32).to_bits(),
1037            "a min of +0.0 must be rewritten to -0.0 when -0.0 is present"
1038        );
1039
1040        let max_ties = [-0.0f32, 0.0, -1.0];
1041        let (_, max_val) = minmax_finite_scalar(&max_ties);
1042        assert_eq!(
1043            max_val.to_bits(),
1044            0.0f32.to_bits(),
1045            "max must take +0.0 when both zero signs are present"
1046        );
1047
1048        let min_ties = [0.0f32, -0.0, 1.0];
1049        let (min_val, _) = minmax_finite_scalar(&min_ties);
1050        assert_eq!(
1051            min_val.to_bits(),
1052            (-0.0f32).to_bits(),
1053            "min must take -0.0 when both zero signs are present"
1054        );
1055
1056        // A bound that is zero with only one sign available keeps that sign.
1057        let (only_neg_min, only_neg_max) = minmax_finite_scalar(&[-0.0f32, -1.0]);
1058        assert_eq!(only_neg_max.to_bits(), (-0.0f32).to_bits());
1059        assert_eq!(only_neg_min.to_bits(), (-1.0f32).to_bits());
1060        let (only_pos_min, only_pos_max) = minmax_finite_scalar(&[0.0f32, 1.0]);
1061        assert_eq!(only_pos_min.to_bits(), 0.0f32.to_bits());
1062        assert_eq!(only_pos_max.to_bits(), 1.0f32.to_bits());
1063
1064        // Non-zero bounds are returned untouched, and an all-nonfinite input keeps the
1065        // identity pair rather than being rewritten by the sign pass.
1066        assert_eq!(
1067            minmax_finite_scalar(&[]),
1068            (f32::INFINITY, f32::NEG_INFINITY)
1069        );
1070        assert_eq!(
1071            minmax_finite_scalar(&[f32::NAN, f32::INFINITY]),
1072            (f32::INFINITY, f32::NEG_INFINITY)
1073        );
1074    }
1075
1076    // FP-034: NEON SDOT vs scalar parity for INT8 dot product.
1077    // Gated on dotprod: SDOT is FEAT_DotProd, not baseline NEON.
1078    #[test]
1079    fn test_i8_neon_scalar_parity() {
1080        #[cfg(target_arch = "aarch64")]
1081        {
1082            if !super::super::SimdConfig::detect().dotprod_enabled {
1083                eprintln!("skipping SDOT parity test: dotprod not available");
1084                return;
1085            }
1086        }
1087        #[cfg(target_arch = "aarch64")]
1088        for dim in [7usize, 16, 64, 128, 384, 768] {
1089            let a_q = QuantizedVector::from_f32(&gen_vec(dim, 200 + dim as u64));
1090            let b_q = QuantizedVector::from_f32(&gen_vec(dim, 300 + dim as u64));
1091
1092            // SAFETY: dotprod confirmed above; slices have equal length from from_f32.
1093            let neon = unsafe { dot_product_i8_neon_unrolled(&a_q.data, &b_q.data) };
1094            let scalar: f32 = a_q
1095                .data
1096                .iter()
1097                .zip(b_q.data.iter())
1098                .map(|(&x, &y)| x as i32 * y as i32)
1099                .sum::<i32>() as f32;
1100
1101            let diff = (neon - scalar).abs();
1102            assert!(
1103                diff <= 1.0,
1104                "NEON vs scalar i8 dot product dim={dim}: neon={neon} scalar={scalar} diff={diff}"
1105            );
1106        }
1107    }
1108
1109    // FP-034: AVX2 vs scalar parity for INT8 dot product.
1110    #[test]
1111    fn test_i8_avx2_scalar_parity() {
1112        #[cfg(target_arch = "x86_64")]
1113        if std::arch::is_x86_feature_detected!("avx2") {
1114            for dim in [7usize, 16, 64, 128, 384, 768] {
1115                let a_q = QuantizedVector::from_f32(&gen_vec(dim, 400 + dim as u64));
1116                let b_q = QuantizedVector::from_f32(&gen_vec(dim, 500 + dim as u64));
1117
1118                // SAFETY: AVX2 verified by is_x86_feature_detected! above; slices have equal length.
1119                let avx2 = unsafe { dot_product_i8_avx2_unrolled(&a_q.data, &b_q.data) };
1120                let scalar: f32 = a_q
1121                    .data
1122                    .iter()
1123                    .zip(b_q.data.iter())
1124                    .map(|(&x, &y)| x as i32 * y as i32)
1125                    .sum::<i32>() as f32;
1126
1127                let diff = (avx2 - scalar).abs();
1128                assert!(
1129                    diff <= 1.0,
1130                    "AVX2 vs scalar i8 dot product dim={dim}: avx2={avx2} scalar={scalar} diff={diff}"
1131                );
1132            }
1133        }
1134    }
1135
1136    // FP-034: AVX-512 VNNI vs scalar parity for INT8 dot product.
1137    // Compiled unconditionally on x86_64 (the kernel is no longer feature-gated);
1138    // runs only when the host advertises AVX-512F+BW+VNNI, matching the kernel's
1139    // safety contract and SimdConfig::avx512vnni_enabled.
1140    #[cfg(target_arch = "x86_64")]
1141    #[test]
1142    fn test_i8_avx512vnni_scalar_parity() {
1143        if std::arch::is_x86_feature_detected!("avx512f")
1144            && std::arch::is_x86_feature_detected!("avx512bw")
1145            && std::arch::is_x86_feature_detected!("avx512vnni")
1146        {
1147            for dim in [7usize, 16, 64, 128, 384, 768] {
1148                let a_q = QuantizedVector::from_f32(&gen_vec(dim, 600 + dim as u64));
1149                let b_q = QuantizedVector::from_f32(&gen_vec(dim, 700 + dim as u64));
1150
1151                // SAFETY: AVX-512F+BW+VNNI verified above; slices have equal length from from_f32.
1152                let vnni = unsafe { dot_product_i8_avx512vnni(&a_q.data, &b_q.data) };
1153                let scalar: f32 = a_q
1154                    .data
1155                    .iter()
1156                    .zip(b_q.data.iter())
1157                    .map(|(&x, &y)| x as i32 * y as i32)
1158                    .sum::<i32>() as f32;
1159
1160                let diff = (vnni - scalar).abs();
1161                assert!(
1162                    diff <= 1.0,
1163                    "VNNI vs scalar i8 dot product dim={dim}: vnni={vnni} scalar={scalar} diff={diff}"
1164                );
1165            }
1166        }
1167    }
1168}