use half::bf16;
pub const Q_MAX: i32 = 127;
#[track_caller]
pub(crate) fn checked_len(context: &str, lhs: usize, rhs: usize, expr: &str) -> usize {
let len = lhs.checked_mul(rhs);
assert!(len.is_some(), "{context}: {expr} overflow ({lhs} * {rhs})");
len.unwrap_or(0)
}
#[derive(Debug, Clone, PartialEq)]
pub struct QuantizedInt8 {
pub q: Vec<i8>,
pub scales: Vec<f32>,
pub n: usize,
pub k: usize,
}
impl QuantizedInt8 {
#[must_use]
pub fn weight_bytes(&self) -> Vec<u8> {
self.q.iter().map(|&v| v as u8).collect()
}
#[must_use]
pub fn scale_bytes(&self) -> Vec<u8> {
self.scales.iter().flat_map(|&s| s.to_le_bytes()).collect()
}
}
#[inline]
#[must_use]
fn round_ties_even_f32(x: f32) -> f32 {
x.round_ties_even()
}
#[must_use]
fn quantize_row(row: &[f32]) -> (Vec<i8>, f32) {
let max_abs = row.iter().fold(0.0f32, |m, &w| m.max(w.abs()));
if max_abs == 0.0 {
return (vec![0i8; row.len()], 1.0);
}
let scale = max_abs / Q_MAX as f32;
let q: Vec<i8> = row
.iter()
.map(|&w| {
let r = round_ties_even_f32(w / scale);
r.clamp(-(Q_MAX as f32), Q_MAX as f32) as i32 as i8
})
.collect();
(q, scale)
}
#[must_use]
pub fn quantize_int8_f32(weights: &[f32], n: usize, k: usize) -> QuantizedInt8 {
let len = checked_len("quantize_int8_f32", n, k, "n*k");
assert_eq!(
weights.len(),
len,
"quantize_int8_f32: weights len {} != n*k {}",
weights.len(),
len
);
let mut q = Vec::with_capacity(len);
let mut scales = Vec::with_capacity(n);
for o in 0..n {
let row = &weights[o * k..(o + 1) * k];
let (q_row, scale) = quantize_row(row);
q.extend_from_slice(&q_row);
scales.push(scale);
}
QuantizedInt8 { q, scales, n, k }
}
#[must_use]
pub fn quantize_int8_f32_searched(
weights: &[f32],
n: usize,
k: usize,
importance: Option<&[f64]>,
) -> QuantizedInt8 {
let len = checked_len("quantize_int8_f32_searched", n, k, "n*k");
assert_eq!(
weights.len(),
len,
"quantize_int8_f32_searched: weights len {} != n*k {}",
weights.len(),
len
);
if let Some(imp) = importance {
assert_eq!(
imp.len(),
k,
"quantize_int8_f32_searched: importance len {} != k {k}",
imp.len()
);
}
let mut q = vec![0i8; len];
let mut scales = vec![0.0f32; n];
{
use rayon::prelude::*;
q.par_chunks_mut(k)
.zip(scales.par_iter_mut())
.zip(weights.par_chunks(k))
.for_each(|((q_row, scale_slot), row)| {
let max_abs = row.iter().fold(0.0f32, |m, &w| m.max(w.abs()));
if max_abs == 0.0 {
*scale_slot = 1.0;
return;
}
let mut best_scale = max_abs / Q_MAX as f32;
let mut best_err = f64::INFINITY;
for i in 0..=super::int4::CLIP_SEARCH_STEPS {
let scale = super::int4::clip_fraction(i) * max_abs / Q_MAX as f32;
if scale <= 0.0 || !scale.is_finite() {
continue;
}
let mut err = 0.0f64;
for (c, &w) in row.iter().enumerate() {
let qv =
round_ties_even_f32(w / scale).clamp(-(Q_MAX as f32), Q_MAX as f32);
let d = f64::from(w - qv * scale);
err += importance.map_or(1.0, |imp| imp[c]) * d * d;
}
if err < best_err {
best_err = err;
best_scale = scale;
}
}
for (slot, &w) in q_row.iter_mut().zip(row.iter()) {
*slot = round_ties_even_f32(w / best_scale).clamp(-(Q_MAX as f32), Q_MAX as f32)
as i32 as i8;
}
*scale_slot = best_scale;
});
}
QuantizedInt8 { q, scales, n, k }
}
#[must_use]
pub fn quantize_int8_bf16(weights: &[bf16], n: usize, k: usize) -> QuantizedInt8 {
let len = checked_len("quantize_int8_bf16", n, k, "n*k");
assert_eq!(
weights.len(),
len,
"quantize_int8_bf16: weights len {} != n*k {}",
weights.len(),
len
);
let widened: Vec<f32> = weights.iter().map(|&w| w.to_f32()).collect();
quantize_int8_f32(&widened, n, k)
}
#[must_use]
pub fn dequantize_int8(q: &QuantizedInt8) -> Vec<f32> {
let len = checked_len("dequantize_int8", q.n, q.k, "n*k");
assert_eq!(
q.q.len(),
len,
"dequantize_int8: q len {} != n*k {}",
q.q.len(),
len
);
assert_eq!(
q.scales.len(),
q.n,
"dequantize_int8: scales len {} != n {}",
q.scales.len(),
q.n
);
let mut out = Vec::with_capacity(len);
for o in 0..q.n {
let s = q.scales[o];
for &v in &q.q[o * q.k..(o + 1) * q.k] {
out.push(s * f32::from(v));
}
}
out
}
#[derive(Debug, Clone, PartialEq)]
pub struct QuantizedU8Activation {
pub q: Vec<u8>,
pub scale: f32,
pub zero_point: i32,
}
#[must_use]
pub fn quantize_activation_u8(x: &[f32]) -> QuantizedU8Activation {
let mut min = 0.0f32;
let mut max = 0.0f32;
for &v in x {
if v < min {
min = v;
}
if v > max {
max = v;
}
}
let range = max - min;
let scale = if range > 0.0 { range / 255.0 } else { 1.0 };
let zp_f = round_ties_even_f32(-min / scale);
let zero_point = zp_f.clamp(0.0, 255.0) as i32;
let q: Vec<u8> = x
.iter()
.map(|&v| {
let r = round_ties_even_f32(v / scale) as i32 + zero_point;
r.clamp(0, 255) as u8
})
.collect();
QuantizedU8Activation {
q,
scale,
zero_point,
}
}
#[must_use]
pub fn dequantize_activation_u8(a: &QuantizedU8Activation) -> Vec<f32> {
a.q.iter()
.map(|&v| a.scale * (i32::from(v) - a.zero_point) as f32)
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn quantizes_known_row_exactly() {
let q = quantize_int8_f32(&[127.0, -127.0, 0.0, 64.0], 1, 4);
assert_eq!(q.n, 1);
assert_eq!(q.k, 4);
assert_eq!(q.scales, vec![1.0]);
assert_eq!(q.q, vec![127i8, -127, 0, 64]);
}
#[test]
fn scale_is_max_abs_over_127() {
let q = quantize_int8_f32(&[254.0, -2.0, 100.0], 1, 3);
assert!((q.scales[0] - 2.0).abs() < 1e-9);
assert_eq!(q.q, vec![127i8, -1, 50]);
}
#[test]
fn per_output_channel_scales_are_independent() {
let w = [127.0f32, 0.0, -64.0, 254.0, -254.0, 0.0];
let q = quantize_int8_f32(&w, 2, 3);
assert_eq!(q.scales.len(), 2);
assert!((q.scales[0] - 1.0).abs() < 1e-9);
assert!((q.scales[1] - 2.0).abs() < 1e-9);
assert_eq!(&q.q[0..3], &[127i8, 0, -64]);
assert_eq!(&q.q[3..6], &[127i8, -127, 0]);
}
#[test]
fn all_zero_row_gets_unit_scale_no_nan() {
let q = quantize_int8_f32(&[0.0, 0.0, -0.0, 0.0], 1, 4);
assert_eq!(q.scales, vec![1.0]);
assert!(q.scales[0].is_finite());
assert_eq!(q.q, vec![0i8; 4]);
}
#[test]
fn round_ties_to_even_at_half() {
let q = quantize_int8_f32(&[127.0, 0.5, 1.5, 2.5], 1, 4);
assert!((q.scales[0] - 1.0).abs() < 1e-9);
assert_eq!(q.q, vec![127i8, 0, 2, 2]);
}
#[test]
fn signed_weight_ties_round_to_even_on_both_sides_of_zero() {
let q = quantize_int8_f32(&[-127.0, -2.5, -1.5, -0.5, 0.5, 1.5, 2.5, 127.0], 1, 8);
assert_eq!(q.scales, vec![1.0]);
assert_eq!(q.q, vec![-127i8, -2, -2, 0, 0, 2, 2, 127]);
let again = quantize_int8_f32(&[-127.0, -2.5, -1.5, -0.5, 0.5, 1.5, 2.5, 127.0], 1, 8);
assert_eq!(q, again, "weight quant must be byte-identical across runs");
}
#[test]
fn negative_value_never_reaches_minus_128() {
let q = quantize_int8_f32(&[-1000.0, 1000.0], 1, 2);
assert_eq!(q.q, vec![-127i8, 127]);
for &v in &q.q {
assert!(v >= -Q_MAX as i8 && v <= Q_MAX as i8);
}
}
#[test]
fn dequant_roundtrips_representable_values() {
let q = quantize_int8_f32(&[254.0, -2.0, 100.0], 1, 3);
let d = dequantize_int8(&q);
assert_eq!(d, vec![254.0, -2.0, 100.0]);
}
#[test]
fn bf16_path_matches_f32_path_on_exact_values() {
let vals = [1.0f32, -2.0, 0.5, 64.0, -64.0, 0.0];
let bf: Vec<bf16> = vals.iter().map(|&v| bf16::from_f32(v)).collect();
let qb = quantize_int8_bf16(&bf, 2, 3);
let qf = quantize_int8_f32(&vals, 2, 3);
assert_eq!(qb, qf);
}
#[test]
fn weight_and_scale_bytes_are_writer_ready() {
let q = quantize_int8_f32(&[127.0, -127.0, 254.0, -254.0], 2, 2);
let wb = q.weight_bytes();
assert_eq!(wb.len(), 4);
assert_eq!(wb[1] as i8, -127i8);
let sb = q.scale_bytes();
assert_eq!(sb.len(), 2 * 4); let s0 = f32::from_le_bytes([sb[0], sb[1], sb[2], sb[3]]);
assert!((s0 - 1.0).abs() < 1e-9);
}
#[test]
#[should_panic(expected = "weights len")]
fn rejects_shape_mismatch() {
let _ = quantize_int8_f32(&[1.0, 2.0, 3.0], 2, 3);
}
#[test]
#[should_panic(expected = "quantize_int8_f32: n*k overflow")]
fn quantize_int8_f32_rejects_shape_product_overflow() {
let _ = quantize_int8_f32(&[], usize::MAX, 2);
}
#[test]
#[should_panic(expected = "quantize_int8_bf16: n*k overflow")]
fn quantize_int8_bf16_rejects_shape_product_overflow() {
let _ = quantize_int8_bf16(&[], usize::MAX, 2);
}
#[test]
#[should_panic(expected = "dequantize_int8: n*k overflow")]
fn dequantize_int8_rejects_shape_product_overflow() {
let q = QuantizedInt8 {
q: Vec::new(),
scales: Vec::new(),
n: usize::MAX,
k: 2,
};
let _ = dequantize_int8(&q);
}
#[test]
fn quant_clamp_bounds_i32_accumulator_at_k_6848() {
const K: usize = 6848;
let a = vec![Q_MAX as i8; K];
let b = vec![Q_MAX as i8; K];
let mut acc: i64 = 0;
for (&x, &w) in a.iter().zip(b.iter()) {
acc += i64::from(x) * i64::from(w);
}
assert_eq!(acc, (K as i64) * 127 * 127);
assert_eq!(acc, 110_451_392);
assert!(acc < i64::from(i32::MAX), "S8S8 K=6848 must fit i32");
let u8s8: i64 = (K as i64) * 255 * 127;
assert_eq!(u8s8, 221_772_480);
assert!(u8s8 < i64::from(i32::MAX), "U8S8 K=6848 must fit i32");
}
#[test]
fn activation_u8_quant_covers_zero_in_range() {
let x = [-1.0f32, 0.0, 1.0, 3.0];
let a = quantize_activation_u8(&x);
let scale = 4.0f32 / 255.0;
assert!((a.scale - scale).abs() < 1e-7);
let zp = round_ties_even_f32(1.0 / scale) as i32;
assert_eq!(a.zero_point, zp);
let d = dequantize_activation_u8(&a);
for (orig, deq) in x.iter().zip(d.iter()) {
assert!((orig - deq).abs() <= a.scale, "{orig} vs {deq}");
}
}
#[test]
fn activation_u8_all_positive_keeps_min_at_zero() {
let x = [1.0f32, 2.0, 255.0];
let a = quantize_activation_u8(&x);
assert_eq!(a.zero_point, 0);
assert!((a.scale - (255.0 / 255.0)).abs() < 1e-7);
assert_eq!(a.q, vec![1u8, 2, 255]);
}
#[test]
fn activation_u8_constant_vector_is_lossless() {
let x = [5.0f32; 6];
let a = quantize_activation_u8(&x);
let d = dequantize_activation_u8(&a);
for &v in &d {
assert!((v - 5.0).abs() <= a.scale + 1e-6);
}
}
#[test]
fn activation_u8_empty_is_unit_scale_no_panic() {
let a = quantize_activation_u8(&[]);
assert!(a.scale.is_finite());
assert_eq!(a.q.len(), 0);
assert_eq!(a.zero_point, 0);
}
#[test]
fn activation_u8_ties_round_to_even_and_are_byte_identical() {
let x = [-127.0f32, 128.0, -2.5, -1.5, -0.5, 0.5, 1.5, 2.5];
let a = quantize_activation_u8(&x);
assert_eq!(a.scale, 1.0);
assert_eq!(a.zero_point, 127);
assert_eq!(a.q, vec![0u8, 255, 125, 125, 127, 127, 129, 129]);
let again = quantize_activation_u8(&x);
assert_eq!(
a, again,
"activation quant must be byte-identical across runs"
);
}
#[test]
fn activation_u8_values_stay_in_byte_range() {
let x = [-1000.0f32, 1000.0, 0.0, -0.0001, 0.0001];
let a = quantize_activation_u8(&x);
for &v in &a.q {
let _ = v;
}
assert!((0..=255).contains(&a.zero_point));
}
fn int8_objective(q: &QuantizedInt8, w: &[f32], importance: Option<&[f64]>) -> f64 {
let deq = dequantize_int8(q);
let mut acc = 0.0f64;
for o in 0..q.n {
for c in 0..q.k {
let d = f64::from(w[o * q.k + c] - deq[o * q.k + c]);
acc += importance.map_or(1.0, |imp| imp[c]) * d * d;
}
}
acc
}
fn int8_lcg(n: usize, seed: u64) -> Vec<f32> {
let mut state = seed | 1;
(0..n)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
((state >> 33) as f32 / (1u64 << 31) as f32) - 1.0
})
.collect()
}
#[test]
fn int8_search_never_loses_to_rtn() {
for (n, k, seed) in [(4usize, 96usize, 5u64), (7, 33, 17), (1, 128, 71)] {
let w = int8_lcg(n * k, seed);
let rtn = quantize_int8_f32(&w, n, k);
let searched = quantize_int8_f32_searched(&w, n, k, None);
assert!(
int8_objective(&searched, &w, None)
<= int8_objective(&rtn, &w, None) * (1.0 + 1e-12),
"shape [{n},{k}] searched must not lose to RTN"
);
let importance: Vec<f64> = (0..k).map(|i| 0.001 + (i % 11) as f64).collect();
let weighted = quantize_int8_f32_searched(&w, n, k, Some(&importance));
assert!(
int8_objective(&weighted, &w, Some(&importance))
<= int8_objective(&rtn, &w, Some(&importance)) * (1.0 + 1e-12)
);
}
}
#[test]
fn int8_search_beats_min_max_on_a_cold_outlier_channel() {
let k = 64usize;
let outlier = 5.0f32;
let rtn_scale = outlier / Q_MAX as f32;
let mut w: Vec<f32> = (0..k)
.map(|i| 40.5 * rtn_scale * if i % 2 == 0 { 1.0 } else { -1.0 })
.collect();
w[k - 1] = outlier;
let mut importance = vec![1.0f64; k];
importance[k - 1] = 1.0e-9;
let rtn = quantize_int8_f32(&w, 1, k);
let searched = quantize_int8_f32_searched(&w, 1, k, Some(&importance));
assert!(
int8_objective(&searched, &w, Some(&importance))
< int8_objective(&rtn, &w, Some(&importance)) * 0.8,
"searched {} vs rtn {}",
int8_objective(&searched, &w, Some(&importance)),
int8_objective(&rtn, &w, Some(&importance))
);
assert!(searched.scales[0] < rtn.scales[0], "the search must clip");
}
#[test]
fn int8_search_is_deterministic_and_preserves_layout() {
let (n, k) = (5usize, 48usize);
let w = int8_lcg(n * k, 2024);
let importance: Vec<f64> = (0..k).map(|i| 1.0 + i as f64).collect();
let first = quantize_int8_f32_searched(&w, n, k, Some(&importance));
for _ in 0..3 {
let again = quantize_int8_f32_searched(&w, n, k, Some(&importance));
assert_eq!(first.q, again.q);
assert_eq!(first.scales, again.scales);
}
assert_eq!(first.q.len(), n * k);
assert_eq!(first.scales.len(), n);
for &v in &first.q {
assert!((-127..=127).contains(&i32::from(v)));
}
let zeros = quantize_int8_f32_searched(&vec![0.0f32; k], 1, k, None);
assert_eq!(zeros.scales, vec![1.0]);
}
#[test]
#[should_panic(expected = "importance len")]
fn int8_search_rejects_a_mismatched_importance_vector() {
let _ = quantize_int8_f32_searched(&[0.0; 16], 1, 16, Some(&[1.0; 4]));
}
}