#[inline]
pub(crate) fn quantize_block_i8(src: &[f32], out: &mut [i8]) -> f32 {
debug_assert_eq!(src.len(), out.len());
#[cfg(target_arch = "x86_64")]
{
if have_avx512_quant() {
return unsafe { quantize_block_i8_avx512(src, out) };
}
}
quantize_block_i8_scalar(src, out)
}
#[inline]
pub(crate) fn quantize_block_u8_offset(src: &[f32], out: &mut [u8]) -> f32 {
debug_assert_eq!(src.len(), out.len());
#[cfg(target_arch = "x86_64")]
{
if have_avx512_quant() {
return unsafe { quantize_block_u8_offset_avx512(src, out) };
}
}
quantize_block_u8_offset_scalar(src, out)
}
#[inline]
pub(crate) fn quantize_block_i16(src: &[f32], out: &mut [i16]) -> f32 {
debug_assert_eq!(src.len(), out.len());
#[cfg(target_arch = "x86_64")]
{
if have_avx512_quant() {
return unsafe { quantize_block_i16_avx512(src, out) };
}
}
quantize_block_i16_scalar(src, out)
}
#[cfg(target_arch = "x86_64")]
fn have_avx512_quant() -> bool {
use std::sync::OnceLock;
static SUPPORTED: OnceLock<bool> = OnceLock::new();
*SUPPORTED.get_or_init(|| {
std::arch::is_x86_feature_detected!("avx512f")
&& std::arch::is_x86_feature_detected!("avx512bw")
&& std::arch::is_x86_feature_detected!("avx512vl")
})
}
fn max_abs_scalar(src: &[f32]) -> f32 {
src.iter().map(|value| value.abs()).fold(0.0f32, f32::max)
}
fn quantize_block_i8_scalar(src: &[f32], out: &mut [i8]) -> f32 {
let max_abs = max_abs_scalar(src);
if max_abs == 0.0 {
out.fill(0);
return 0.0;
}
let scale = max_abs / 127.0;
let inverse_scale = 127.0 / max_abs;
for (code, &value) in out.iter_mut().zip(src) {
*code = (value * inverse_scale).round().clamp(-127.0, 127.0) as i8;
}
scale
}
fn quantize_block_u8_offset_scalar(src: &[f32], out: &mut [u8]) -> f32 {
let max_abs = max_abs_scalar(src);
if max_abs == 0.0 {
out.fill(128);
return 0.0;
}
let scale = max_abs / 127.0;
let inverse_scale = 127.0 / max_abs;
for (code, &value) in out.iter_mut().zip(src) {
let signed = (value * inverse_scale).round().clamp(-127.0, 127.0) as i8;
*code = (signed as i16 + 128) as u8;
}
scale
}
fn quantize_block_i16_scalar(src: &[f32], out: &mut [i16]) -> f32 {
let max_abs = max_abs_scalar(src);
if max_abs == 0.0 {
out.fill(0);
return 0.0;
}
let scale = max_abs / 32767.0;
let inverse_scale = 32767.0 / max_abs;
for (code, &value) in out.iter_mut().zip(src) {
*code = (value * inverse_scale).round().clamp(-32767.0, 32767.0) as i16;
}
scale
}
#[cfg(target_arch = "x86_64")]
struct MaxAbsReduction {
max_abs: f32,
all_finite: bool,
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn max_abs_avx512(src: &[f32]) -> MaxAbsReduction {
use std::arch::x86_64::*;
unsafe {
let n = src.len();
let ptr = src.as_ptr();
let sign = _mm512_set1_ps(-0.0);
let infinity = _mm512_set1_ps(f32::INFINITY);
let mut acc = _mm512_setzero_ps();
let mut all_finite = true;
let mut i = 0;
while i + 16 <= n {
let value = _mm512_loadu_ps(ptr.add(i));
let magnitude = _mm512_andnot_ps(sign, value);
let finite = _mm512_cmp_ps_mask::<_CMP_LT_OQ>(magnitude, infinity);
all_finite &= finite == 0xFFFF;
acc = _mm512_max_ps(acc, magnitude);
i += 16;
}
let mut max_abs = _mm512_reduce_max_ps(acc);
while i < n {
let magnitude = (*ptr.add(i)).abs();
all_finite &= magnitude.is_finite();
max_abs = max_abs.max(magnitude);
i += 1;
}
MaxAbsReduction { max_abs, all_finite }
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
#[inline]
unsafe fn round_half_away_avx512(
values: std::arch::x86_64::__m512,
) -> std::arch::x86_64::__m512 {
use std::arch::x86_64::*;
unsafe {
let sign = _mm512_set1_ps(-0.0);
let one = _mm512_set1_ps(1.0);
let half = _mm512_set1_ps(0.5);
let truncated = _mm512_roundscale_ps::<0x0B>(values);
let fraction = _mm512_sub_ps(values, truncated);
let abs_fraction = _mm512_andnot_ps(sign, fraction);
let at_or_past_half = _mm512_cmp_ps_mask::<_CMP_GE_OQ>(abs_fraction, half);
let signed_one = _mm512_or_ps(_mm512_and_ps(values, sign), one);
let bump = _mm512_maskz_mov_ps(at_or_past_half, signed_one);
_mm512_add_ps(truncated, bump)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vl")]
unsafe fn quantize_block_i8_avx512(src: &[f32], out: &mut [i8]) -> f32 {
use std::arch::x86_64::*;
unsafe {
let n = src.len();
let reduction = max_abs_avx512(src);
if !reduction.all_finite {
return quantize_block_i8_scalar(src, out);
}
let max_abs = reduction.max_abs;
if max_abs == 0.0 {
out.fill(0);
return 0.0;
}
let scale = max_abs / 127.0;
let inverse_scale = 127.0 / max_abs;
let inverse = _mm512_set1_ps(inverse_scale);
let lower = _mm512_set1_ps(-127.0);
let upper = _mm512_set1_ps(127.0);
let src_ptr = src.as_ptr();
let out_ptr = out.as_mut_ptr();
let mut i = 0;
while i + 16 <= n {
let value = _mm512_loadu_ps(src_ptr.add(i));
let scaled = _mm512_mul_ps(value, inverse);
let rounded = round_half_away_avx512(scaled);
let clamped = _mm512_min_ps(_mm512_max_ps(rounded, lower), upper);
let as_i32 = _mm512_cvttps_epi32(clamped);
let as_i8 = _mm512_cvtepi32_epi8(as_i32);
_mm_storeu_si128(out_ptr.add(i).cast(), as_i8);
i += 16;
}
while i < n {
*out_ptr.add(i) = (*src_ptr.add(i) * inverse_scale).round().clamp(-127.0, 127.0) as i8;
i += 1;
}
scale
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vl")]
unsafe fn quantize_block_u8_offset_avx512(src: &[f32], out: &mut [u8]) -> f32 {
use std::arch::x86_64::*;
unsafe {
let n = src.len();
let reduction = max_abs_avx512(src);
if !reduction.all_finite {
return quantize_block_u8_offset_scalar(src, out);
}
let max_abs = reduction.max_abs;
if max_abs == 0.0 {
out.fill(128);
return 0.0;
}
let scale = max_abs / 127.0;
let inverse_scale = 127.0 / max_abs;
let inverse = _mm512_set1_ps(inverse_scale);
let lower = _mm512_set1_ps(-127.0);
let upper = _mm512_set1_ps(127.0);
let offset = _mm512_set1_epi32(128);
let src_ptr = src.as_ptr();
let out_ptr = out.as_mut_ptr();
let mut i = 0;
while i + 16 <= n {
let value = _mm512_loadu_ps(src_ptr.add(i));
let scaled = _mm512_mul_ps(value, inverse);
let rounded = round_half_away_avx512(scaled);
let clamped = _mm512_min_ps(_mm512_max_ps(rounded, lower), upper);
let as_i32 = _mm512_cvttps_epi32(clamped);
let offset_i32 = _mm512_add_epi32(as_i32, offset);
let as_u8 = _mm512_cvtepi32_epi8(offset_i32);
_mm_storeu_si128(out_ptr.add(i).cast(), as_u8);
i += 16;
}
while i < n {
let signed =
(*src_ptr.add(i) * inverse_scale).round().clamp(-127.0, 127.0) as i8;
*out_ptr.add(i) = (signed as i16 + 128) as u8;
i += 1;
}
scale
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f,avx512bw,avx512vl")]
unsafe fn quantize_block_i16_avx512(src: &[f32], out: &mut [i16]) -> f32 {
use std::arch::x86_64::*;
unsafe {
let n = src.len();
let reduction = max_abs_avx512(src);
if !reduction.all_finite {
return quantize_block_i16_scalar(src, out);
}
let max_abs = reduction.max_abs;
if max_abs == 0.0 {
out.fill(0);
return 0.0;
}
let scale = max_abs / 32767.0;
let inverse_scale = 32767.0 / max_abs;
let inverse = _mm512_set1_ps(inverse_scale);
let lower = _mm512_set1_ps(-32767.0);
let upper = _mm512_set1_ps(32767.0);
let src_ptr = src.as_ptr();
let out_ptr = out.as_mut_ptr();
let mut i = 0;
while i + 16 <= n {
let value = _mm512_loadu_ps(src_ptr.add(i));
let scaled = _mm512_mul_ps(value, inverse);
let rounded = round_half_away_avx512(scaled);
let clamped = _mm512_min_ps(_mm512_max_ps(rounded, lower), upper);
let as_i32 = _mm512_cvttps_epi32(clamped);
let as_i16 = _mm512_cvtepi32_epi16(as_i32);
_mm256_storeu_si256(out_ptr.add(i).cast(), as_i16);
i += 16;
}
while i < n {
*out_ptr.add(i) =
(*src_ptr.add(i) * inverse_scale).round().clamp(-32767.0, 32767.0) as i16;
i += 1;
}
scale
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_row(len: usize, seed: u64) -> Vec<f32> {
let mut state = seed.wrapping_mul(0x9E3779B97F4A7C15).wrapping_add(1);
(0..len)
.map(|_| {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
let unit = ((state >> 11) as f64 / (1u64 << 53) as f64) as f32;
let base = (unit - 0.5) * 6.0;
if state.is_multiple_of(97) {
base * 20.0
} else {
base
}
})
.collect()
}
const LENGTHS: [usize; 16] =
[1, 2, 3, 7, 15, 16, 17, 31, 32, 33, 48, 64, 127, 128, 129, 1024];
#[test]
fn quantize_block_i8_simd_matches_scalar_bit_identical() {
for &len in &LENGTHS {
for seed in 0..16u64 {
let row = make_row(len, seed);
let mut simd = vec![0i8; len];
let mut scalar = vec![7i8; len];
let scale_simd = quantize_block_i8(&row, &mut simd);
let scale_scalar = quantize_block_i8_scalar(&row, &mut scalar);
assert_eq!(scale_simd.to_bits(), scale_scalar.to_bits(), "len {len} seed {seed}");
assert_eq!(simd, scalar, "len {len} seed {seed}");
}
}
}
#[test]
fn quantize_block_u8_offset_simd_matches_scalar_bit_identical() {
for &len in &LENGTHS {
for seed in 0..16u64 {
let row = make_row(len, seed);
let mut simd = vec![0u8; len];
let mut scalar = vec![7u8; len];
let scale_simd = quantize_block_u8_offset(&row, &mut simd);
let scale_scalar = quantize_block_u8_offset_scalar(&row, &mut scalar);
assert_eq!(scale_simd.to_bits(), scale_scalar.to_bits(), "len {len} seed {seed}");
assert_eq!(simd, scalar, "len {len} seed {seed}");
}
}
}
#[test]
fn quantize_block_i16_simd_matches_scalar_bit_identical() {
for &len in &LENGTHS {
for seed in 0..16u64 {
let row = make_row(len, seed);
let mut simd = vec![0i16; len];
let mut scalar = vec![7i16; len];
let scale_simd = quantize_block_i16(&row, &mut simd);
let scale_scalar = quantize_block_i16_scalar(&row, &mut scalar);
assert_eq!(scale_simd.to_bits(), scale_scalar.to_bits(), "len {len} seed {seed}");
assert_eq!(simd, scalar, "len {len} seed {seed}");
}
}
}
#[test]
fn quantize_edge_cases_bit_identical() {
let mut rows: Vec<Vec<f32>> = Vec::new();
rows.push(vec![0.0; 40]);
rows.push(vec![-0.0; 40]);
rows.push(vec![5.0; 40]);
rows.push(vec![-5.0; 40]);
rows.push((0..40).map(|i| if i % 2 == 0 { 3.0 } else { -3.0 }).collect());
{
let mut row = vec![0.01f32; 40];
row[19] = 100.0;
row[20] = -100.0;
rows.push(row);
}
{
let max_abs = 127.0f32;
let mut row = vec![0.0f32; 40];
for (i, value) in row.iter_mut().enumerate() {
*value = ((i as f32) - 20.0) + 0.5; }
row[0] = max_abs; rows.push(row);
}
for row in &rows {
let mut i8_simd = vec![0i8; row.len()];
let mut i8_scalar = vec![0i8; row.len()];
assert_eq!(
quantize_block_i8(row, &mut i8_simd).to_bits(),
quantize_block_i8_scalar(row, &mut i8_scalar).to_bits()
);
assert_eq!(i8_simd, i8_scalar);
let mut u8_simd = vec![0u8; row.len()];
let mut u8_scalar = vec![0u8; row.len()];
assert_eq!(
quantize_block_u8_offset(row, &mut u8_simd).to_bits(),
quantize_block_u8_offset_scalar(row, &mut u8_scalar).to_bits()
);
assert_eq!(u8_simd, u8_scalar);
let mut i16_simd = vec![0i16; row.len()];
let mut i16_scalar = vec![0i16; row.len()];
assert_eq!(
quantize_block_i16(row, &mut i16_simd).to_bits(),
quantize_block_i16_scalar(row, &mut i16_scalar).to_bits()
);
assert_eq!(i16_simd, i16_scalar);
}
}
#[test]
fn quantize_non_finite_and_signed_zero_bit_identical() {
let nan = f32::NAN;
let inf = f32::INFINITY;
let mut rows: Vec<Vec<f32>> = Vec::new();
{
let mut row: Vec<f32> = (0..16).map(|i| (i as f32) - 7.5).collect();
row[15] = nan;
rows.push(row);
}
{
let mut row: Vec<f32> = (0..32).map(|i| ((i as f32) - 16.0) * 0.3).collect();
row[31] = nan;
rows.push(row);
}
{
let mut row: Vec<f32> = (0..40).map(|i| (i as f32) - 20.0).collect();
row[0] = nan;
rows.push(row);
}
rows.push(vec![nan; 40]);
rows.push(vec![nan; 16]);
rows.push(vec![nan; 7]);
{
let mut row = vec![1.0f32; 40];
row[5] = inf;
row[6] = -inf;
rows.push(row);
}
rows.push(vec![inf; 16]);
rows.push(vec![-inf; 20]);
rows.push(vec![-0.0f32; 40]);
rows.push((0..40).map(|i| if i % 2 == 0 { -0.0 } else { 0.0 }).collect());
{
let mut row = vec![0.0f32; 33];
row[3] = nan;
row[8] = inf;
row[9] = -inf;
row[10] = -0.0;
row[17] = 12.5;
row[31] = -9.25;
rows.push(row);
}
for row in &rows {
let len = row.len();
let mut i8_simd = vec![0i8; len];
let mut i8_scalar = vec![0i8; len];
let s8_simd = quantize_block_i8(row, &mut i8_simd);
let s8_scalar = quantize_block_i8_scalar(row, &mut i8_scalar);
assert_eq!(s8_simd.to_bits(), s8_scalar.to_bits(), "i8 scale, len {len}");
assert_eq!(i8_simd, i8_scalar, "i8 codes, len {len}");
let mut u8_simd = vec![0u8; len];
let mut u8_scalar = vec![0u8; len];
let su_simd = quantize_block_u8_offset(row, &mut u8_simd);
let su_scalar = quantize_block_u8_offset_scalar(row, &mut u8_scalar);
assert_eq!(su_simd.to_bits(), su_scalar.to_bits(), "u8 scale, len {len}");
assert_eq!(u8_simd, u8_scalar, "u8 codes, len {len}");
let mut i16_simd = vec![0i16; len];
let mut i16_scalar = vec![0i16; len];
let s16_simd = quantize_block_i16(row, &mut i16_simd);
let s16_scalar = quantize_block_i16_scalar(row, &mut i16_scalar);
assert_eq!(s16_simd.to_bits(), s16_scalar.to_bits(), "i16 scale, len {len}");
assert_eq!(i16_simd, i16_scalar, "i16 codes, len {len}");
}
}
}