#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
use std::sync::OnceLock;
use super::simd_config;
#[derive(Debug, Clone, Copy)]
pub struct QuantizationParams {
pub scale: f32,
pub zero_point: i8,
pub min_val: f32,
pub max_val: f32,
}
impl QuantizationParams {
pub fn from_vector(vector: &[f32]) -> Self {
let (mut min_val, mut max_val) = minmax_finite(vector);
if !min_val.is_finite() || !max_val.is_finite() {
min_val = 0.0;
max_val = 0.0;
}
let max_abs = min_val.abs().max(max_val.abs());
let scale = if max_abs > 1e-10 {
127.0 / max_abs
} else {
1.0 };
Self {
scale,
zero_point: 0,
min_val,
max_val,
}
}
}
fn minmax_finite(v: &[f32]) -> (f32, f32) {
#[cfg(target_arch = "x86_64")]
{
if simd_config().avx2_enabled {
return unsafe { minmax_finite_avx2(v) };
}
}
#[cfg(target_arch = "aarch64")]
{
if simd_config().neon_enabled {
return unsafe { minmax_finite_neon(v) };
}
}
minmax_finite_scalar(v)
}
#[inline]
fn pin_zero_signs(v: &[f32], min_val: f32, max_val: f32) -> (f32, f32) {
if min_val != 0.0 && max_val != 0.0 {
return (min_val, max_val);
}
let mut has_negative_zero = false;
let mut has_positive_zero = false;
for &value in v {
has_negative_zero |= value.to_bits() == (-0.0f32).to_bits();
has_positive_zero |= value.to_bits() == 0.0f32.to_bits();
}
let min_val = if min_val == 0.0 {
if has_negative_zero { -0.0 } else { 0.0 }
} else {
min_val
};
let max_val = if max_val == 0.0 {
if has_positive_zero { 0.0 } else { -0.0 }
} else {
max_val
};
(min_val, max_val)
}
fn minmax_finite_scalar(v: &[f32]) -> (f32, f32) {
let mut min_val = f32::INFINITY;
let mut max_val = f32::NEG_INFINITY;
for &x in v {
if x.is_finite() {
min_val = min_val.min(x);
max_val = max_val.max(x);
}
}
pin_zero_signs(v, min_val, max_val)
}
#[cfg(test)]
thread_local! {
static I8_MINMAX_SIMD_HITS: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn minmax_finite_avx2(v: &[f32]) -> (f32, f32) {
#[cfg(test)]
I8_MINMAX_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
let chunks = v.len() / 8;
let inf = _mm256_set1_ps(f32::INFINITY);
let neg_inf = _mm256_set1_ps(f32::NEG_INFINITY);
let sign = _mm256_set1_ps(-0.0);
let mut vmin = inf;
let mut vmax = neg_inf;
for i in 0..chunks {
let x = _mm256_loadu_ps(v.as_ptr().add(i * 8));
let abs = _mm256_andnot_ps(sign, x);
let finite = _mm256_cmp_ps(abs, inf, _CMP_LT_OQ);
vmin = _mm256_min_ps(vmin, _mm256_blendv_ps(inf, x, finite));
vmax = _mm256_max_ps(vmax, _mm256_blendv_ps(neg_inf, x, finite));
}
let mut min_lanes = [0.0f32; 8];
let mut max_lanes = [0.0f32; 8];
_mm256_storeu_ps(min_lanes.as_mut_ptr(), vmin);
_mm256_storeu_ps(max_lanes.as_mut_ptr(), vmax);
let mut min_val = min_lanes.into_iter().fold(f32::INFINITY, f32::min);
let mut max_val = max_lanes.into_iter().fold(f32::NEG_INFINITY, f32::max);
for &x in &v[chunks * 8..] {
if x.is_finite() {
min_val = min_val.min(x);
max_val = max_val.max(x);
}
}
pin_zero_signs(v, min_val, max_val)
}
#[cfg(target_arch = "aarch64")]
unsafe fn minmax_finite_neon(v: &[f32]) -> (f32, f32) {
#[cfg(test)]
I8_MINMAX_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
let chunks = v.len() / 4;
let inf = unsafe { vdupq_n_f32(f32::INFINITY) };
let neg_inf = unsafe { vdupq_n_f32(f32::NEG_INFINITY) };
let mut vmin = inf;
let mut vmax = neg_inf;
for i in 0..chunks {
let x = unsafe { vld1q_f32(v.as_ptr().add(i * 4)) };
unsafe {
let finite = vcaltq_f32(x, inf);
vmin = vminq_f32(vmin, vbslq_f32(finite, x, inf));
vmax = vmaxq_f32(vmax, vbslq_f32(finite, x, neg_inf));
}
}
let (mut min_val, mut max_val) = unsafe { (vminvq_f32(vmin), vmaxvq_f32(vmax)) };
for &x in &v[chunks * 4..] {
if x.is_finite() {
min_val = min_val.min(x);
max_val = max_val.max(x);
}
}
pin_zero_signs(v, min_val, max_val)
}
#[derive(Debug, Clone)]
pub struct QuantizedVector {
data: Vec<i8>,
pub params: QuantizationParams,
pub norm: f32,
}
impl QuantizedVector {
#[inline]
pub fn data(&self) -> &[i8] {
&self.data
}
#[inline]
pub fn len(&self) -> usize {
self.data.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
}
impl QuantizedVector {
pub fn from_f32(vector: &[f32]) -> Self {
let mut params = QuantizationParams::from_vector(vector);
if !params.scale.is_finite() || params.scale == 0.0 {
params.scale = 1.0;
}
let mut norm_sq = 0.0f32;
for &v in vector {
if v.is_finite() {
norm_sq += v * v;
}
}
let norm = norm_sq.sqrt();
let data = quantize_i8(vector, params.scale);
Self { data, params, norm }
}
pub fn to_f32(&self) -> Vec<f32> {
let scale = if self.params.scale.is_finite() && self.params.scale != 0.0 {
self.params.scale
} else {
1.0
};
self.data.iter().map(|&v| v as f32 / scale).collect()
}
#[inline]
pub fn dot_product(&self, other: &QuantizedVector) -> f32 {
dot_product_i8(self, other)
}
#[inline]
pub fn cosine_similarity(&self, other: &QuantizedVector) -> f32 {
cosine_similarity_i8(self, other)
}
}
fn quantize_i8(vector: &[f32], scale: f32) -> Vec<i8> {
#[cfg(target_arch = "x86_64")]
{
if simd_config().avx2_enabled {
return unsafe { quantize_i8_avx2(vector, scale) };
}
}
#[cfg(target_arch = "aarch64")]
{
if simd_config().neon_enabled {
return unsafe { quantize_i8_neon(vector, scale) };
}
}
quantize_i8_scalar(vector, scale)
}
fn quantize_i8_scalar(vector: &[f32], scale: f32) -> Vec<i8> {
vector
.iter()
.map(|&v| quantize_i8_value(v, scale))
.collect()
}
#[inline]
fn quantize_i8_value(value: f32, scale: f32) -> i8 {
if value.is_finite() {
(value * scale).round().clamp(-127.0, 127.0) as i8
} else {
0
}
}
#[cfg(test)]
thread_local! {
static I8_QUANTIZE_SIMD_HITS: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn quantize_i8_avx2(vector: &[f32], scale: f32) -> Vec<i8> {
#[cfg(test)]
I8_QUANTIZE_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
let mut data = vec![0i8; vector.len()];
let chunks = vector.len() / 8;
let scale_scalar = scale;
let scale = _mm256_set1_ps(scale_scalar);
let inf = _mm256_set1_ps(f32::INFINITY);
let sign = _mm256_set1_ps(-0.0);
let low = _mm256_set1_ps(-127.0);
let high = _mm256_set1_ps(127.0);
let half = _mm256_set1_ps(0.5);
let negative_half = _mm256_set1_ps(-0.5);
let one = _mm256_set1_epi32(1);
let negative_one = _mm256_set1_epi32(-1);
for i in 0..chunks {
let base = i * 8;
let input = _mm256_loadu_ps(vector.as_ptr().add(base));
let abs = _mm256_andnot_ps(sign, input);
let finite = _mm256_cmp_ps(abs, inf, _CMP_LT_OQ);
let values = _mm256_and_ps(input, finite);
let scaled = _mm256_mul_ps(values, scale);
let clamped = _mm256_min_ps(_mm256_max_ps(scaled, low), high);
let truncated = _mm256_cvttps_epi32(clamped);
let fraction = _mm256_sub_ps(clamped, _mm256_cvtepi32_ps(truncated));
let round_up = _mm256_castps_si256(_mm256_cmp_ps(fraction, half, _CMP_GE_OQ));
let round_down = _mm256_castps_si256(_mm256_cmp_ps(fraction, negative_half, _CMP_LE_OQ));
let rounded = _mm256_add_epi32(
_mm256_add_epi32(truncated, _mm256_and_si256(round_up, one)),
_mm256_and_si256(round_down, negative_one),
);
let mut lanes = [0i32; 8];
_mm256_storeu_si256(lanes.as_mut_ptr().cast::<__m256i>(), rounded);
for (offset, lane) in lanes.into_iter().enumerate() {
data[base + offset] = lane as i8;
}
}
for i in chunks * 8..vector.len() {
data[i] = quantize_i8_value(vector[i], scale_scalar);
}
data
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn quantize_i8_neon(vector: &[f32], scale: f32) -> Vec<i8> {
#[cfg(test)]
I8_QUANTIZE_SIMD_HITS.with(|hits| hits.set(hits.get() + 1));
let mut data = vec![0i8; vector.len()];
let chunks = vector.len() / 4;
let scale_vector = vdupq_n_f32(scale);
let inf = vdupq_n_f32(f32::INFINITY);
let zero = vdupq_n_f32(0.0);
let low = vdupq_n_f32(-127.0);
let high = vdupq_n_f32(127.0);
for i in 0..chunks {
let base = i * 4;
let input = vld1q_f32(vector.as_ptr().add(base));
let finite = vcaltq_f32(input, inf);
let values = vbslq_f32(finite, input, zero);
let scaled = vmulq_f32(values, scale_vector);
let clamped = vminq_f32(vmaxq_f32(scaled, low), high);
let rounded = vcvtaq_s32_f32(clamped);
let mut lanes = [0i32; 4];
vst1q_s32(lanes.as_mut_ptr(), rounded);
for (offset, lane) in lanes.into_iter().enumerate() {
data[base + offset] = lane as i8;
}
}
for i in chunks * 4..vector.len() {
data[i] = quantize_i8_value(vector[i], scale);
}
data
}
#[inline]
pub fn dot_product_i8(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
debug_assert!(a.data.iter().all(|&v| v != -128i8));
debug_assert!(b.data.iter().all(|&v| v != -128i8));
if a.data.len() != b.data.len() {
return 0.0;
}
let denom = a.params.scale * b.params.scale;
if denom == 0.0 || !denom.is_finite() {
return 0.0;
}
dot_product_i8_dispatch(&a.data, &b.data) / denom
}
#[inline]
pub(crate) fn dot_product_i8_trusted(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
if a.data.len() != b.data.len() {
return 0.0;
}
let denom = a.params.scale * b.params.scale;
if denom == 0.0 || !denom.is_finite() {
return 0.0;
}
debug_assert!(a.data.iter().all(|&v| v != i8::MIN));
debug_assert!(b.data.iter().all(|&v| v != i8::MIN));
dot_product_i8_dispatch(&a.data, &b.data) / denom
}
#[inline]
pub fn cosine_similarity_i8(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
let denom = a.norm * b.norm;
if denom == 0.0 || !denom.is_finite() {
return 0.0;
}
dot_product_i8(a, b) / denom
}
#[inline]
pub(crate) fn cosine_similarity_i8_trusted(a: &QuantizedVector, b: &QuantizedVector) -> f32 {
let denom = a.norm * b.norm;
if denom == 0.0 || !denom.is_finite() {
return 0.0;
}
dot_product_i8_trusted(a, b) / denom
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "dotprod")]
unsafe fn dot_product_i8_neon_unrolled(a: &[i8], b: &[i8]) -> f32 {
const SIMD_WIDTH: usize = 16;
const UNROLL: usize = 4;
const CHUNK_SIZE: usize = SIMD_WIDTH * UNROLL;
const PREFETCH_DISTANCE: usize = CHUNK_SIZE;
let n = a.len();
debug_assert_eq!(n, b.len());
let chunks = n / CHUNK_SIZE;
let mut sum0 = vdupq_n_s32(0);
let mut sum1 = vdupq_n_s32(0);
let mut sum2 = vdupq_n_s32(0);
let mut sum3 = vdupq_n_s32(0);
for i in 0..chunks {
let base = i * CHUNK_SIZE;
let next_base = base + PREFETCH_DISTANCE;
if next_base + CHUNK_SIZE <= n {
core::arch::asm!(
"prfm pldl1keep, [{ptr}]",
ptr = in(reg) a.as_ptr().add(next_base),
options(nostack, readonly, preserves_flags)
);
core::arch::asm!(
"prfm pldl1keep, [{ptr}]",
ptr = in(reg) b.as_ptr().add(next_base),
options(nostack, readonly, preserves_flags)
);
}
let a0 = vld1q_s8(a.as_ptr().add(base));
let b0 = vld1q_s8(b.as_ptr().add(base));
let a1 = vld1q_s8(a.as_ptr().add(base + SIMD_WIDTH));
let b1 = vld1q_s8(b.as_ptr().add(base + SIMD_WIDTH));
let a2 = vld1q_s8(a.as_ptr().add(base + SIMD_WIDTH * 2));
let b2 = vld1q_s8(b.as_ptr().add(base + SIMD_WIDTH * 2));
let a3 = vld1q_s8(a.as_ptr().add(base + SIMD_WIDTH * 3));
let b3 = vld1q_s8(b.as_ptr().add(base + SIMD_WIDTH * 3));
core::arch::asm!(
"sdot {s0:v}.4s, {a0:v}.16b, {b0:v}.16b",
"sdot {s1:v}.4s, {a1:v}.16b, {b1:v}.16b",
"sdot {s2:v}.4s, {a2:v}.16b, {b2:v}.16b",
"sdot {s3:v}.4s, {a3:v}.16b, {b3:v}.16b",
s0 = inout(vreg) sum0,
a0 = in(vreg) a0,
b0 = in(vreg) b0,
s1 = inout(vreg) sum1,
a1 = in(vreg) a1,
b1 = in(vreg) b1,
s2 = inout(vreg) sum2,
a2 = in(vreg) a2,
b2 = in(vreg) b2,
s3 = inout(vreg) sum3,
a3 = in(vreg) a3,
b3 = in(vreg) b3,
options(nomem, nostack, preserves_flags)
);
}
let sum01 = vaddq_s32(sum0, sum1);
let sum23 = vaddq_s32(sum2, sum3);
let mut sum_vec = vaddq_s32(sum01, sum23);
let tail_start = chunks * CHUNK_SIZE;
let tail_chunks = (n - tail_start) / SIMD_WIDTH;
for j in 0..tail_chunks {
let base = tail_start + j * SIMD_WIDTH;
let at = vld1q_s8(a.as_ptr().add(base));
let bt = vld1q_s8(b.as_ptr().add(base));
core::arch::asm!(
"sdot {acc:v}.4s, {a:v}.16b, {b:v}.16b",
acc = inout(vreg) sum_vec,
a = in(vreg) at,
b = in(vreg) bt,
options(nomem, nostack, preserves_flags)
);
}
let sum = vaddvq_s32(sum_vec);
let remainder_start = tail_start + tail_chunks * SIMD_WIDTH;
let remainder: i32 = a[remainder_start..]
.iter()
.zip(b[remainder_start..].iter())
.map(|(&x, &y)| x as i32 * y as i32)
.sum();
(sum + remainder) as f32
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f", enable = "avx512bw")]
#[inline]
unsafe fn mm512_sign_epi8(b: __m512i, a: __m512i) -> __m512i {
let zero = _mm512_setzero_si512();
let neg_b = _mm512_sub_epi8(zero, b);
let mask_neg = _mm512_cmplt_epi8_mask(a, zero);
let mask_zero = _mm512_cmpeq_epi8_mask(a, zero);
let result = _mm512_mask_blend_epi8(mask_neg, b, neg_b);
_mm512_mask_blend_epi8(mask_zero, result, zero)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f", enable = "avx512vnni", enable = "avx512bw")]
unsafe fn dot_product_i8_avx512vnni(a: &[i8], b: &[i8]) -> f32 {
const SIMD_WIDTH: usize = 64; const UNROLL: usize = 4;
const CHUNK_SIZE: usize = SIMD_WIDTH * UNROLL;
let n = a.len();
debug_assert_eq!(n, b.len());
debug_assert!(a.iter().all(|&v| v != i8::MIN));
debug_assert!(b.iter().all(|&v| v != i8::MIN));
let chunks = n / CHUNK_SIZE;
let mut sum0 = _mm512_setzero_si512();
let mut sum1 = _mm512_setzero_si512();
let mut sum2 = _mm512_setzero_si512();
let mut sum3 = _mm512_setzero_si512();
for i in 0..chunks {
let base = i * CHUNK_SIZE;
let a0 = _mm512_loadu_si512(a.as_ptr().add(base) as *const __m512i);
let b0 = _mm512_loadu_si512(b.as_ptr().add(base) as *const __m512i);
let a0_abs = _mm512_abs_epi8(a0);
let b0_signed = mm512_sign_epi8(b0, a0);
sum0 = _mm512_dpbusd_epi32(sum0, a0_abs, b0_signed);
let a1 = _mm512_loadu_si512(a.as_ptr().add(base + SIMD_WIDTH) as *const __m512i);
let b1 = _mm512_loadu_si512(b.as_ptr().add(base + SIMD_WIDTH) as *const __m512i);
let a1_abs = _mm512_abs_epi8(a1);
let b1_signed = mm512_sign_epi8(b1, a1);
sum1 = _mm512_dpbusd_epi32(sum1, a1_abs, b1_signed);
let a2 = _mm512_loadu_si512(a.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m512i);
let b2 = _mm512_loadu_si512(b.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m512i);
let a2_abs = _mm512_abs_epi8(a2);
let b2_signed = mm512_sign_epi8(b2, a2);
sum2 = _mm512_dpbusd_epi32(sum2, a2_abs, b2_signed);
let a3 = _mm512_loadu_si512(a.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m512i);
let b3 = _mm512_loadu_si512(b.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m512i);
let a3_abs = _mm512_abs_epi8(a3);
let b3_signed = mm512_sign_epi8(b3, a3);
sum3 = _mm512_dpbusd_epi32(sum3, a3_abs, b3_signed);
}
let sum01 = _mm512_add_epi32(sum0, sum1);
let sum23 = _mm512_add_epi32(sum2, sum3);
let sum_vec = _mm512_add_epi32(sum01, sum23);
let sum = _mm512_reduce_add_epi32(sum_vec);
let remainder_start = chunks * CHUNK_SIZE;
let remainder: i32 = a[remainder_start..]
.iter()
.zip(b[remainder_start..].iter())
.map(|(&x, &y)| x as i32 * y as i32)
.sum();
(sum + remainder) as f32
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_product_i8_avx2_unrolled(a: &[i8], b: &[i8]) -> f32 {
const SIMD_WIDTH: usize = 32;
const UNROLL: usize = 4;
const CHUNK_SIZE: usize = SIMD_WIDTH * UNROLL;
const PREFETCH_DISTANCE: usize = CHUNK_SIZE;
let n = a.len();
debug_assert_eq!(n, b.len());
debug_assert!(a.iter().all(|&v| v != i8::MIN));
debug_assert!(b.iter().all(|&v| v != i8::MIN));
let chunks = n / CHUNK_SIZE;
let mut sum0 = _mm256_setzero_si256();
let mut sum1 = _mm256_setzero_si256();
let mut sum2 = _mm256_setzero_si256();
let mut sum3 = _mm256_setzero_si256();
let ones = _mm256_set1_epi16(1);
for i in 0..chunks {
let base = i * CHUNK_SIZE;
let next_base = base + PREFETCH_DISTANCE;
if next_base + CHUNK_SIZE <= n {
_mm_prefetch(a.as_ptr().add(next_base), _MM_HINT_T0);
_mm_prefetch(b.as_ptr().add(next_base), _MM_HINT_T0);
}
let a0 = _mm256_loadu_si256(a.as_ptr().add(base) as *const __m256i);
let b0 = _mm256_loadu_si256(b.as_ptr().add(base) as *const __m256i);
let prod0 = _mm256_maddubs_epi16(_mm256_abs_epi8(a0), _mm256_sign_epi8(b0, a0));
let prod0_32 = _mm256_madd_epi16(prod0, ones);
sum0 = _mm256_add_epi32(sum0, prod0_32);
let a1 = _mm256_loadu_si256(a.as_ptr().add(base + SIMD_WIDTH) as *const __m256i);
let b1 = _mm256_loadu_si256(b.as_ptr().add(base + SIMD_WIDTH) as *const __m256i);
let prod1 = _mm256_maddubs_epi16(_mm256_abs_epi8(a1), _mm256_sign_epi8(b1, a1));
let prod1_32 = _mm256_madd_epi16(prod1, ones);
sum1 = _mm256_add_epi32(sum1, prod1_32);
let a2 = _mm256_loadu_si256(a.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m256i);
let b2 = _mm256_loadu_si256(b.as_ptr().add(base + SIMD_WIDTH * 2) as *const __m256i);
let prod2 = _mm256_maddubs_epi16(_mm256_abs_epi8(a2), _mm256_sign_epi8(b2, a2));
let prod2_32 = _mm256_madd_epi16(prod2, ones);
sum2 = _mm256_add_epi32(sum2, prod2_32);
let a3 = _mm256_loadu_si256(a.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m256i);
let b3 = _mm256_loadu_si256(b.as_ptr().add(base + SIMD_WIDTH * 3) as *const __m256i);
let prod3 = _mm256_maddubs_epi16(_mm256_abs_epi8(a3), _mm256_sign_epi8(b3, a3));
let prod3_32 = _mm256_madd_epi16(prod3, ones);
sum3 = _mm256_add_epi32(sum3, prod3_32);
}
let sum01 = _mm256_add_epi32(sum0, sum1);
let sum23 = _mm256_add_epi32(sum2, sum3);
let sum_vec = _mm256_add_epi32(sum01, sum23);
let sum128_lo = _mm256_castsi256_si128(sum_vec);
let sum128_hi = _mm256_extracti128_si256(sum_vec, 1);
let sum128 = _mm_add_epi32(sum128_lo, sum128_hi);
let sum64 = _mm_add_epi32(sum128, _mm_srli_si128(sum128, 8));
let sum32 = _mm_add_epi32(sum64, _mm_srli_si128(sum64, 4));
let sum = _mm_cvtsi128_si32(sum32);
let remainder_start = chunks * CHUNK_SIZE;
let remainder: i32 = a[remainder_start..]
.iter()
.zip(b[remainder_start..].iter())
.map(|(&x, &y)| x as i32 * y as i32)
.sum();
(sum + remainder) as f32
}
pub type I8DotKernel = fn(&[i8], &[i8]) -> f32;
static I8_DOT_KERNEL: OnceLock<I8DotKernel> = OnceLock::new();
#[inline]
pub fn resolved_i8_dot_kernel() -> I8DotKernel {
*I8_DOT_KERNEL.get_or_init(resolve_i8_dot_kernel)
}
fn resolve_i8_dot_kernel() -> I8DotKernel {
let config = simd_config();
#[cfg(target_arch = "aarch64")]
{
if config.neon_enabled && config.dotprod_enabled {
return dot_product_i8_neon_kernel;
}
}
#[cfg(target_arch = "x86_64")]
{
if config.avx512vnni_enabled {
return dot_product_i8_avx512vnni_kernel;
}
if config.avx2_enabled {
return dot_product_i8_avx2_kernel;
}
}
dot_product_i8_scalar_kernel
}
#[cfg(target_arch = "aarch64")]
fn dot_product_i8_neon_kernel(a: &[i8], b: &[i8]) -> f32 {
unsafe { dot_product_i8_neon_unrolled(a, b) }
}
#[cfg(target_arch = "x86_64")]
fn dot_product_i8_avx512vnni_kernel(a: &[i8], b: &[i8]) -> f32 {
debug_assert!(a.iter().all(|&v| v != i8::MIN));
debug_assert!(b.iter().all(|&v| v != i8::MIN));
unsafe { dot_product_i8_avx512vnni(a, b) }
}
#[cfg(target_arch = "x86_64")]
fn dot_product_i8_avx2_kernel(a: &[i8], b: &[i8]) -> f32 {
debug_assert!(a.iter().all(|&v| v != i8::MIN));
debug_assert!(b.iter().all(|&v| v != i8::MIN));
unsafe { dot_product_i8_avx2_unrolled(a, b) }
}
fn dot_product_i8_scalar_kernel(a: &[i8], b: &[i8]) -> f32 {
a.iter()
.zip(b.iter())
.map(|(&x, &y)| x as i32 * y as i32)
.sum::<i32>() as f32
}
#[inline]
fn dot_product_i8_dispatch(a: &[i8], b: &[i8]) -> f32 {
resolved_i8_dot_kernel()(a, b)
}
#[inline]
pub fn dot_product_i8_raw(a: &[i8], b: &[i8]) -> f32 {
if a.len() != b.len() {
return 0.0;
}
debug_assert!(
a.iter().all(|&v| v != -128i8),
"dot_product_i8_raw: slice a contains -128, violating the [-127, 127] SIMD invariant"
);
debug_assert!(
b.iter().all(|&v| v != -128i8),
"dot_product_i8_raw: slice b contains -128, violating the [-127, 127] SIMD invariant"
);
dot_product_i8_dispatch(a, b)
}
#[cfg(test)]
mod simd_parity_tests {
use super::*;
fn gen_vec(dim: usize, seed: u64) -> Vec<f32> {
let mut state = seed ^ ((dim as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15));
(0..dim)
.map(|i| {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407)
.wrapping_add(i as u64);
let unit = ((state >> 32) as u32) as f32 / u32::MAX as f32;
unit * 2.0 - 1.0
})
.collect()
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[test]
fn test_i8_quantize_explicit_simd_matches_scalar_and_is_dispatched() {
#[cfg(target_arch = "x86_64")]
if !std::arch::is_x86_feature_detected!("avx2") {
return;
}
for dim in [0usize, 1, 3, 4, 7, 8, 9, 31, 32, 33, 383, 384, 385] {
let mut input = gen_vec(dim, 900 + dim as u64);
if dim > 0 {
input[0] = f32::NAN;
}
if dim > 1 {
input[1] = f32::INFINITY;
}
if dim > 2 {
input[2] = f32::NEG_INFINITY;
}
if dim > 3 {
input[3] = 0.25;
}
if dim > 4 {
input[4] = -0.25;
}
if dim > 5 {
input[5] = f32::from_bits(0.25f32.to_bits() - 1);
}
if dim > 6 {
input[6] = f32::from_bits(0.25f32.to_bits() + 1);
}
let scalar = quantize_i8_scalar(&input, 2.0);
#[cfg(target_arch = "aarch64")]
let simd = unsafe { quantize_i8_neon(&input, 2.0) };
#[cfg(target_arch = "x86_64")]
let simd = unsafe { quantize_i8_avx2(&input, 2.0) };
assert_eq!(simd, scalar, "explicit SIMD mismatch at dim={dim}");
}
let input = gen_vec(385, 1_063);
let before = I8_QUANTIZE_SIMD_HITS.with(std::cell::Cell::get);
let quantized = QuantizedVector::from_f32(&input);
let after = I8_QUANTIZE_SIMD_HITS.with(std::cell::Cell::get);
assert_eq!(
after,
before + 1,
"QuantizedVector::from_f32 did not execute its explicit SIMD quantizer"
);
assert_eq!(
quantized.data,
quantize_i8_scalar(&input, quantized.params.scale)
);
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
#[test]
fn test_finite_minmax_explicit_simd_matches_scalar() {
#[cfg(target_arch = "x86_64")]
if !std::arch::is_x86_feature_detected!("avx2") {
return;
}
for dim in [0usize, 1, 3, 4, 7, 8, 9, 31, 32, 33, 383, 384, 385] {
let mut input = gen_vec(dim, 1_100 + dim as u64);
if dim > 0 {
input[0] = f32::NAN;
}
if dim > 1 {
input[1] = f32::INFINITY;
}
if dim > 2 {
input[2] = f32::NEG_INFINITY;
}
if dim > 3 {
input[3] = -0.0;
}
if dim > 4 {
input[4] = 0.0;
}
let scalar = minmax_finite_scalar(&input);
#[cfg(target_arch = "aarch64")]
let simd = unsafe { minmax_finite_neon(&input) };
#[cfg(target_arch = "x86_64")]
let simd = unsafe { minmax_finite_avx2(&input) };
assert_eq!(
(simd.0.to_bits(), simd.1.to_bits()),
(scalar.0.to_bits(), scalar.1.to_bits()),
"finite min/max mismatch at dim={dim}"
);
}
let mut min_zero_input = vec![1.0f32; 16];
min_zero_input[0] = -0.0;
min_zero_input[8] = 0.0;
let mut max_zero_input = vec![-1.0f32; 16];
max_zero_input[0] = 0.0;
max_zero_input[8] = -0.0;
for input in [&min_zero_input, &max_zero_input] {
let scalar = minmax_finite_scalar(input);
#[cfg(target_arch = "aarch64")]
let simd = unsafe { minmax_finite_neon(input) };
#[cfg(target_arch = "x86_64")]
let simd = unsafe { minmax_finite_avx2(input) };
assert_eq!(
(simd.0.to_bits(), simd.1.to_bits()),
(scalar.0.to_bits(), scalar.1.to_bits()),
"finite min/max signed-zero mismatch"
);
}
let input = gen_vec(385, 1_063);
let before = I8_MINMAX_SIMD_HITS.with(std::cell::Cell::get);
let params = QuantizationParams::from_vector(&input);
let after = I8_MINMAX_SIMD_HITS.with(std::cell::Cell::get);
assert_eq!(
after,
before + 1,
"QuantizationParams::from_vector did not execute its explicit SIMD reducer"
);
let scalar = minmax_finite_scalar(&input);
assert_eq!(
(params.min_val.to_bits(), params.max_val.to_bits()),
(scalar.0.to_bits(), scalar.1.to_bits())
);
}
#[test]
fn test_minmax_finite_pins_zero_signs() {
let both_signs = [-0.0f32, 0.0, -1.0];
assert_eq!(
pin_zero_signs(&both_signs, -1.0, -0.0).1.to_bits(),
0.0f32.to_bits(),
"a max of -0.0 must be rewritten to +0.0 when +0.0 is present"
);
assert_eq!(
pin_zero_signs(&both_signs, 0.0, -1.0).0.to_bits(),
(-0.0f32).to_bits(),
"a min of +0.0 must be rewritten to -0.0 when -0.0 is present"
);
let max_ties = [-0.0f32, 0.0, -1.0];
let (_, max_val) = minmax_finite_scalar(&max_ties);
assert_eq!(
max_val.to_bits(),
0.0f32.to_bits(),
"max must take +0.0 when both zero signs are present"
);
let min_ties = [0.0f32, -0.0, 1.0];
let (min_val, _) = minmax_finite_scalar(&min_ties);
assert_eq!(
min_val.to_bits(),
(-0.0f32).to_bits(),
"min must take -0.0 when both zero signs are present"
);
let (only_neg_min, only_neg_max) = minmax_finite_scalar(&[-0.0f32, -1.0]);
assert_eq!(only_neg_max.to_bits(), (-0.0f32).to_bits());
assert_eq!(only_neg_min.to_bits(), (-1.0f32).to_bits());
let (only_pos_min, only_pos_max) = minmax_finite_scalar(&[0.0f32, 1.0]);
assert_eq!(only_pos_min.to_bits(), 0.0f32.to_bits());
assert_eq!(only_pos_max.to_bits(), 1.0f32.to_bits());
assert_eq!(
minmax_finite_scalar(&[]),
(f32::INFINITY, f32::NEG_INFINITY)
);
assert_eq!(
minmax_finite_scalar(&[f32::NAN, f32::INFINITY]),
(f32::INFINITY, f32::NEG_INFINITY)
);
}
#[test]
fn test_i8_neon_scalar_parity() {
#[cfg(target_arch = "aarch64")]
{
if !super::super::SimdConfig::detect().dotprod_enabled {
eprintln!("skipping SDOT parity test: dotprod not available");
return;
}
}
#[cfg(target_arch = "aarch64")]
for dim in [7usize, 16, 64, 128, 384, 768] {
let a_q = QuantizedVector::from_f32(&gen_vec(dim, 200 + dim as u64));
let b_q = QuantizedVector::from_f32(&gen_vec(dim, 300 + dim as u64));
let neon = unsafe { dot_product_i8_neon_unrolled(&a_q.data, &b_q.data) };
let scalar: f32 = a_q
.data
.iter()
.zip(b_q.data.iter())
.map(|(&x, &y)| x as i32 * y as i32)
.sum::<i32>() as f32;
let diff = (neon - scalar).abs();
assert!(
diff <= 1.0,
"NEON vs scalar i8 dot product dim={dim}: neon={neon} scalar={scalar} diff={diff}"
);
}
}
#[test]
fn test_i8_avx2_scalar_parity() {
#[cfg(target_arch = "x86_64")]
if std::arch::is_x86_feature_detected!("avx2") {
for dim in [7usize, 16, 64, 128, 384, 768] {
let a_q = QuantizedVector::from_f32(&gen_vec(dim, 400 + dim as u64));
let b_q = QuantizedVector::from_f32(&gen_vec(dim, 500 + dim as u64));
let avx2 = unsafe { dot_product_i8_avx2_unrolled(&a_q.data, &b_q.data) };
let scalar: f32 = a_q
.data
.iter()
.zip(b_q.data.iter())
.map(|(&x, &y)| x as i32 * y as i32)
.sum::<i32>() as f32;
let diff = (avx2 - scalar).abs();
assert!(
diff <= 1.0,
"AVX2 vs scalar i8 dot product dim={dim}: avx2={avx2} scalar={scalar} diff={diff}"
);
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn test_i8_avx512vnni_scalar_parity() {
if std::arch::is_x86_feature_detected!("avx512f")
&& std::arch::is_x86_feature_detected!("avx512bw")
&& std::arch::is_x86_feature_detected!("avx512vnni")
{
for dim in [7usize, 16, 64, 128, 384, 768] {
let a_q = QuantizedVector::from_f32(&gen_vec(dim, 600 + dim as u64));
let b_q = QuantizedVector::from_f32(&gen_vec(dim, 700 + dim as u64));
let vnni = unsafe { dot_product_i8_avx512vnni(&a_q.data, &b_q.data) };
let scalar: f32 = a_q
.data
.iter()
.zip(b_q.data.iter())
.map(|(&x, &y)| x as i32 * y as i32)
.sum::<i32>() as f32;
let diff = (vnni - scalar).abs();
assert!(
diff <= 1.0,
"VNNI vs scalar i8 dot product dim={dim}: vnni={vnni} scalar={scalar} diff={diff}"
);
}
}
}
}