1#[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#[derive(Debug, Clone, Copy)]
21pub struct QuantizationParams {
22 pub scale: f32,
24 pub zero_point: i8,
26 pub min_val: f32,
28 pub max_val: f32,
30}
31
32impl QuantizationParams {
33 pub fn from_vector(vector: &[f32]) -> Self {
37 let (mut min_val, mut max_val) = minmax_finite(vector);
39
40 if !min_val.is_finite() || !max_val.is_finite() {
42 min_val = 0.0;
43 max_val = 0.0;
44 }
45
46 let max_abs = min_val.abs().max(max_val.abs());
48
49 let scale = if max_abs > 1e-10 {
51 127.0 / max_abs
52 } else {
53 1.0 };
55
56 Self {
57 scale,
58 zero_point: 0,
59 min_val,
60 max_val,
61 }
62 }
63}
64
65fn minmax_finite(v: &[f32]) -> (f32, f32) {
76 #[cfg(target_arch = "x86_64")]
77 {
78 if simd_config().avx2_enabled {
79 return unsafe { minmax_finite_avx2(v) };
81 }
82 }
83 #[cfg(target_arch = "aarch64")]
84 {
85 if simd_config().neon_enabled {
86 return unsafe { minmax_finite_neon(v) };
88 }
89 }
90 minmax_finite_scalar(v)
91}
92
93#[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 let x = unsafe { vld1q_f32(v.as_ptr().add(i * 4)) };
196 unsafe {
197 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#[derive(Debug, Clone)]
218pub struct QuantizedVector {
219 data: Vec<i8>,
222 pub params: QuantizationParams,
224 pub norm: f32,
226}
227
228impl QuantizedVector {
229 #[inline]
231 pub fn data(&self) -> &[i8] {
232 &self.data
233 }
234
235 #[inline]
237 pub fn len(&self) -> usize {
238 self.data.len()
239 }
240
241 #[inline]
243 pub fn is_empty(&self) -> bool {
244 self.data.is_empty()
245 }
246}
247
248impl QuantizedVector {
249 pub fn from_f32(vector: &[f32]) -> Self {
251 let mut params = QuantizationParams::from_vector(vector);
252
253 if !params.scale.is_finite() || params.scale == 0.0 {
255 params.scale = 1.0;
256 }
257
258 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 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 #[inline]
287 pub fn dot_product(&self, other: &QuantizedVector) -> f32 {
288 dot_product_i8(self, other)
289 }
290
291 #[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 return unsafe { quantize_i8_avx2(vector, scale) };
304 }
305 }
306 #[cfg(target_arch = "aarch64")]
307 {
308 if simd_config().neon_enabled {
309 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#[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#[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#[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#[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#[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 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 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 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#[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 let mask_neg = _mm512_cmplt_epi8_mask(a, zero);
599 let mask_zero = _mm512_cmpeq_epi8_mask(a, zero);
601 let result = _mm512_mask_blend_epi8(mask_neg, b, neg_b);
603 _mm512_mask_blend_epi8(mask_zero, result, zero)
605}
606
607#[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; 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 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 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 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 let sum = _mm512_reduce_add_epi32(sum_vec);
667
668 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#[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 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 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 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 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 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 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 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 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 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 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
769pub type I8DotKernel = fn(&[i8], &[i8]) -> f32;
775
776static I8_DOT_KERNEL: OnceLock<I8DotKernel> = OnceLock::new();
777
778#[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 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 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 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 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#[inline]
840fn dot_product_i8_dispatch(a: &[i8], b: &[i8]) -> f32 {
841 resolved_i8_dot_kernel()(a, b)
842}
843
844#[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 let simd = unsafe { quantize_i8_neon(&input, 2.0) };
918 #[cfg(target_arch = "x86_64")]
919 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 let simd = unsafe { minmax_finite_neon(&input) };
969 #[cfg(target_arch = "x86_64")]
970 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 let simd = unsafe { minmax_finite_neon(input) };
990 #[cfg(target_arch = "x86_64")]
991 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 #[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 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 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 #[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 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 #[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 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 #[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 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}