use crate::GpuOptimError;
use scirs2_core::ndarray::{Array1, Array2, Axis, Zip};
use scirs2_core::random::{Rng, RngExt};
pub const E4M3_MAX_NORMAL: f64 = 448.0;
pub const E5M2_MAX_NORMAL: f64 = 57344.0;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum QuantScheme {
Symmetric,
Affine,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum IntDtype {
Int8,
Int4,
}
impl IntDtype {
pub fn bits(self) -> u32 {
match self {
IntDtype::Int8 => 8,
IntDtype::Int4 => 4,
}
}
pub fn from_bits(bits: u32) -> Result<Self, GpuOptimError> {
match bits {
8 => Ok(IntDtype::Int8),
4 => Ok(IntDtype::Int4),
other => Err(GpuOptimError::UnsupportedOperation(format!(
"unsupported integer quantization width: {other} bits (expected 4 or 8)"
))),
}
}
pub fn q_range(self, scheme: QuantScheme) -> (i32, i32) {
match (self, scheme) {
(IntDtype::Int8, QuantScheme::Symmetric) => (-127, 127),
(IntDtype::Int8, QuantScheme::Affine) => (-128, 127),
(IntDtype::Int4, QuantScheme::Symmetric) => (-7, 7),
(IntDtype::Int4, QuantScheme::Affine) => (-8, 7),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Fp8Format {
E4M3,
E5M2,
}
impl Fp8Format {
pub fn mantissa_bits(self) -> i32 {
match self {
Fp8Format::E4M3 => 3,
Fp8Format::E5M2 => 2,
}
}
pub fn exponent_bias(self) -> i32 {
match self {
Fp8Format::E4M3 => 7,
Fp8Format::E5M2 => 15,
}
}
pub fn max_normal(self) -> f64 {
match self {
Fp8Format::E4M3 => E4M3_MAX_NORMAL,
Fp8Format::E5M2 => E5M2_MAX_NORMAL,
}
}
pub fn min_normal(self) -> f64 {
pow2(1 - self.exponent_bias())
}
pub fn min_subnormal(self) -> f64 {
pow2(1 - self.exponent_bias() - self.mantissa_bits())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RoundingMode {
Nearest,
Stochastic,
}
impl RoundingMode {
fn round_to_int(self, x: f64, rng: &mut impl Rng) -> f64 {
match self {
RoundingMode::Nearest => x.round_ties_even(),
RoundingMode::Stochastic => {
let lower = x.floor();
let frac = x - lower;
let draw: f64 = rng.random();
if draw < frac {
lower + 1.0
} else {
lower
}
}
}
}
}
fn pow2(exp: i32) -> f64 {
2.0_f64.powi(exp)
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct QuantParams {
pub scale: f64,
pub zero_point: i32,
pub qmin: i32,
pub qmax: i32,
}
impl QuantParams {
pub fn from_minmax(
min_v: f64,
max_v: f64,
dtype: IntDtype,
scheme: QuantScheme,
) -> Result<Self, GpuOptimError> {
if !min_v.is_finite() || !max_v.is_finite() || max_v < min_v {
return Err(GpuOptimError::InvalidState(format!(
"invalid calibration interval: [{min_v}, {max_v}]"
)));
}
match scheme {
QuantScheme::Symmetric => Self::from_absmax(min_v.abs().max(max_v.abs()), dtype),
QuantScheme::Affine => {
let (qmin, qmax) = dtype.q_range(QuantScheme::Affine);
let rmin = min_v.min(0.0);
let rmax = max_v.max(0.0);
let range = rmax - rmin;
if range <= 0.0 {
return Err(GpuOptimError::InvalidState(
"degenerate (zero-range) calibration: real min == max == 0".to_string(),
));
}
let scale = range / f64::from(qmax - qmin);
if !scale.is_finite() || scale <= 0.0 {
return Err(GpuOptimError::InvalidState(format!(
"degenerate affine scale derived from interval [{min_v}, {max_v}]"
)));
}
let zp = (f64::from(qmin) - rmin / scale)
.round_ties_even()
.clamp(f64::from(qmin), f64::from(qmax));
Ok(Self {
scale,
zero_point: zp as i32,
qmin,
qmax,
})
}
}
}
pub fn from_absmax(absmax: f64, dtype: IntDtype) -> Result<Self, GpuOptimError> {
let (qmin, qmax) = dtype.q_range(QuantScheme::Symmetric);
if !absmax.is_finite() || absmax <= 0.0 {
return Err(GpuOptimError::InvalidState(
"degenerate (zero-range) calibration: absmax must be finite and > 0".to_string(),
));
}
let scale = absmax / f64::from(qmax);
if !scale.is_finite() || scale <= 0.0 {
return Err(GpuOptimError::InvalidState(
"degenerate symmetric scale".to_string(),
));
}
Ok(Self {
scale,
zero_point: 0,
qmin,
qmax,
})
}
pub fn per_tensor_minmax(
x: &Array1<f64>,
dtype: IntDtype,
scheme: QuantScheme,
) -> Result<Self, GpuOptimError> {
let (min_v, max_v) = finite_min_max(x.iter().copied())?;
Self::from_minmax(min_v, max_v, dtype, scheme)
}
pub fn symmetric_from_absmax(x: &Array1<f64>, dtype: IntDtype) -> Result<Self, GpuOptimError> {
if x.is_empty() {
return Err(GpuOptimError::InvalidState(
"cannot calibrate an empty tensor".to_string(),
));
}
let mut absmax = 0.0_f64;
for &v in x.iter() {
if !v.is_finite() {
return Err(GpuOptimError::InvalidState(
"non-finite value in calibration tensor".to_string(),
));
}
absmax = absmax.max(v.abs());
}
Self::from_absmax(absmax, dtype)
}
pub fn per_tensor_percentile(
x: &Array1<f64>,
dtype: IntDtype,
scheme: QuantScheme,
clip_fraction: f64,
) -> Result<Self, GpuOptimError> {
if !(0.0..0.5).contains(&clip_fraction) {
return Err(GpuOptimError::InvalidState(format!(
"clip_fraction must lie in [0, 0.5), got {clip_fraction}"
)));
}
if x.is_empty() {
return Err(GpuOptimError::InvalidState(
"cannot calibrate an empty tensor".to_string(),
));
}
let mut sorted: Vec<f64> = Vec::with_capacity(x.len());
for &v in x.iter() {
if !v.is_finite() {
return Err(GpuOptimError::InvalidState(
"non-finite value in calibration tensor".to_string(),
));
}
sorted.push(v);
}
sorted.sort_by(|a, b| a.total_cmp(b));
let last = sorted.len() - 1;
let lo_idx = (clip_fraction * last as f64).floor() as usize;
let hi_idx = ((1.0 - clip_fraction) * last as f64).ceil() as usize;
let min_v = sorted[lo_idx.min(last)];
let max_v = sorted[hi_idx.min(last)];
Self::from_minmax(min_v, max_v, dtype, scheme)
}
pub fn quantize(&self, x: f64, mode: RoundingMode, rng: &mut impl Rng) -> i32 {
let scaled = x / self.scale + f64::from(self.zero_point);
let rounded = mode
.round_to_int(scaled, rng)
.clamp(f64::from(self.qmin), f64::from(self.qmax));
rounded as i32
}
pub fn dequantize(&self, q: i32) -> f64 {
f64::from(q - self.zero_point) * self.scale
}
pub fn fake_quant_scalar(&self, x: f64, mode: RoundingMode, rng: &mut impl Rng) -> f64 {
self.dequantize(self.quantize(x, mode, rng))
}
pub fn real_min(&self) -> f64 {
self.dequantize(self.qmin)
}
pub fn real_max(&self) -> f64 {
self.dequantize(self.qmax)
}
}
fn finite_min_max(values: impl Iterator<Item = f64>) -> Result<(f64, f64), GpuOptimError> {
let mut min_v = f64::INFINITY;
let mut max_v = f64::NEG_INFINITY;
let mut count = 0_usize;
for v in values {
if !v.is_finite() {
return Err(GpuOptimError::InvalidState(
"non-finite value in calibration tensor".to_string(),
));
}
min_v = min_v.min(v);
max_v = max_v.max(v);
count += 1;
}
if count == 0 {
return Err(GpuOptimError::InvalidState(
"cannot calibrate an empty tensor".to_string(),
));
}
Ok((min_v, max_v))
}
pub fn fake_quant_int(
x: &Array1<f64>,
params: &QuantParams,
mode: RoundingMode,
rng: &mut impl Rng,
) -> Array1<f64> {
let mut out = Vec::with_capacity(x.len());
for &v in x.iter() {
out.push(params.fake_quant_scalar(v, mode, rng));
}
Array1::from_vec(out)
}
pub fn per_channel_params(
x: &Array2<f64>,
axis: usize,
dtype: IntDtype,
scheme: QuantScheme,
) -> Result<Vec<QuantParams>, GpuOptimError> {
if axis > 1 {
return Err(GpuOptimError::InvalidState(format!(
"per-channel axis must be 0 or 1, got {axis}"
)));
}
let n_channels = x.shape()[axis];
let mut params = Vec::with_capacity(n_channels);
for channel in 0..n_channels {
let lane = x.index_axis(Axis(axis), channel);
let (min_v, max_v) = finite_min_max(lane.iter().copied())?;
params.push(QuantParams::from_minmax(min_v, max_v, dtype, scheme)?);
}
Ok(params)
}
pub fn fake_quant_int_per_channel(
x: &Array2<f64>,
params: &[QuantParams],
axis: usize,
mode: RoundingMode,
rng: &mut impl Rng,
) -> Result<Array2<f64>, GpuOptimError> {
if axis > 1 {
return Err(GpuOptimError::InvalidState(format!(
"per-channel axis must be 0 or 1, got {axis}"
)));
}
let n_channels = x.shape()[axis];
if params.len() != n_channels {
return Err(GpuOptimError::DimensionMismatch {
expected: vec![n_channels],
actual: vec![params.len()],
});
}
let (n_rows, n_cols) = (x.shape()[0], x.shape()[1]);
let mut out = Array2::<f64>::zeros((n_rows, n_cols));
for r in 0..n_rows {
for c in 0..n_cols {
let channel = if axis == 0 { r } else { c };
let p = ¶ms[channel];
out[[r, c]] = p.fake_quant_scalar(x[[r, c]], mode, rng);
}
}
Ok(out)
}
fn fp8_quantize_magnitude(
a: f64,
mantissa_bits: i32,
exponent_bias: i32,
max_normal: f64,
mode: RoundingMode,
rng: &mut impl Rng,
) -> f64 {
if a == 0.0 {
return 0.0;
}
let emin = 1 - exponent_bias;
let e = ((a.to_bits() >> 52) & 0x7ff) as i32 - 1023;
let step_exp = e.max(emin) - mantissa_bits;
let step = pow2(step_exp);
let ratio = a / step;
let rounded = mode.round_to_int(ratio, rng);
let magnitude = rounded * step;
if magnitude > max_normal {
max_normal
} else {
magnitude
}
}
pub fn fake_quant_fp8(
x: &Array1<f64>,
format: Fp8Format,
mode: RoundingMode,
rng: &mut impl Rng,
) -> Array1<f64> {
let mantissa_bits = format.mantissa_bits();
let exponent_bias = format.exponent_bias();
let max_normal = format.max_normal();
let mut out = Vec::with_capacity(x.len());
for &v in x.iter() {
let q = if v.is_nan() {
f64::NAN
} else if v.is_infinite() {
max_normal.copysign(v)
} else {
let magnitude = fp8_quantize_magnitude(
v.abs(),
mantissa_bits,
exponent_bias,
max_normal,
mode,
rng,
);
magnitude.copysign(v)
};
out.push(q);
}
Array1::from_vec(out)
}
pub fn fake_quant_backward(
grad: &Array1<f64>,
x: &Array1<f64>,
qmin_real: f64,
qmax_real: f64,
) -> Result<Array1<f64>, GpuOptimError> {
if grad.len() != x.len() {
return Err(GpuOptimError::DimensionMismatch {
expected: vec![x.len()],
actual: vec![grad.len()],
});
}
let (lo, hi) = if qmin_real <= qmax_real {
(qmin_real, qmax_real)
} else {
(qmax_real, qmin_real)
};
let mut out = Vec::with_capacity(grad.len());
for (&g, &v) in grad.iter().zip(x.iter()) {
if v >= lo && v <= hi {
out.push(g);
} else {
out.push(0.0);
}
}
Ok(Array1::from_vec(out))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum QuantTarget {
Int(IntDtype),
Fp8(Fp8Format),
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct QatConfig {
pub target: QuantTarget,
pub scheme: QuantScheme,
pub rounding: RoundingMode,
pub lr: f64,
pub beta1: f64,
pub beta2: f64,
pub eps: f64,
pub weight_decay: f64,
}
impl QatConfig {
pub fn new(target: QuantTarget, scheme: QuantScheme, rounding: RoundingMode, lr: f64) -> Self {
Self {
target,
scheme,
rounding,
lr,
beta1: 0.9,
beta2: 0.999,
eps: 1e-8,
weight_decay: 0.0,
}
}
}
#[derive(Debug, Clone)]
pub struct QatOptimizer {
config: QatConfig,
first_moment: Array1<f64>,
second_moment: Array1<f64>,
step_count: u64,
quantized: Array1<f64>,
int_params: Option<QuantParams>,
}
impl QatOptimizer {
pub fn new(
master: &Array1<f64>,
config: QatConfig,
rng: &mut impl Rng,
) -> Result<Self, GpuOptimError> {
if master.is_empty() {
return Err(GpuOptimError::InvalidState(
"cannot construct a QatOptimizer over empty master weights".to_string(),
));
}
let n = master.len();
let mut optimizer = Self {
config,
first_moment: Array1::zeros(n),
second_moment: Array1::zeros(n),
step_count: 0,
quantized: Array1::zeros(n),
int_params: None,
};
optimizer.requantize(master, rng)?;
Ok(optimizer)
}
pub fn quantized_weights(&self) -> &Array1<f64> {
&self.quantized
}
pub fn quant_params(&self) -> Option<&QuantParams> {
self.int_params.as_ref()
}
pub fn step_count(&self) -> u64 {
self.step_count
}
pub fn step(
&mut self,
master: &mut Array1<f64>,
grad: &Array1<f64>,
rng: &mut impl Rng,
) -> Result<(), GpuOptimError> {
if master.len() != grad.len() {
return Err(GpuOptimError::DimensionMismatch {
expected: vec![master.len()],
actual: vec![grad.len()],
});
}
if master.len() != self.first_moment.len() {
return Err(GpuOptimError::DimensionMismatch {
expected: vec![self.first_moment.len()],
actual: vec![master.len()],
});
}
self.step_count += 1;
let t = self.step_count as i32;
let beta1 = self.config.beta1;
let beta2 = self.config.beta2;
let lr = self.config.lr;
let eps = self.config.eps;
let weight_decay = self.config.weight_decay;
let bias_correction1 = 1.0 - beta1.powi(t);
let bias_correction2 = 1.0 - beta2.powi(t);
Zip::from(&mut *master)
.and(grad)
.and(&mut self.first_moment)
.and(&mut self.second_moment)
.for_each(|weight, &g, m, v| {
*m = beta1 * *m + (1.0 - beta1) * g;
*v = beta2 * *v + (1.0 - beta2) * g * g;
let m_hat = *m / bias_correction1;
let v_hat = *v / bias_correction2;
if weight_decay != 0.0 {
*weight -= lr * weight_decay * *weight;
}
*weight -= lr * m_hat / (v_hat.sqrt() + eps);
});
self.requantize(master, rng)?;
Ok(())
}
fn requantize(
&mut self,
master: &Array1<f64>,
rng: &mut impl Rng,
) -> Result<(), GpuOptimError> {
match self.config.target {
QuantTarget::Int(dtype) => {
let params = QuantParams::per_tensor_minmax(master, dtype, self.config.scheme)?;
self.quantized = fake_quant_int(master, ¶ms, self.config.rounding, rng);
self.int_params = Some(params);
}
QuantTarget::Fp8(format) => {
self.quantized = fake_quant_fp8(master, format, self.config.rounding, rng);
self.int_params = None;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::random::Random;
const EPS: f64 = 1e-9;
fn seeded(seed: u64) -> Random<scirs2_core::random::rngs::StdRng> {
Random::seed(seed)
}
#[test]
fn int8_grid_values_quantize_to_themselves() {
let params = QuantParams::from_absmax(12.7, IntDtype::Int8).expect("calibrate");
assert!((params.scale - 0.1).abs() < EPS);
assert_eq!(params.zero_point, 0);
let mut rng = seeded(1);
for k in -120..=120 {
let on_grid = k as f64 * params.scale;
let round_trip = params.fake_quant_scalar(on_grid, RoundingMode::Nearest, &mut rng);
assert!(
(round_trip - on_grid).abs() < EPS,
"grid value {on_grid} did not map to itself (got {round_trip})"
);
}
}
#[test]
fn int8_affine_grid_values_quantize_to_themselves() {
let data = Array1::from_vec(vec![-0.3, 1.7, 0.0, 0.9, -0.1]);
let params = QuantParams::per_tensor_minmax(&data, IntDtype::Int8, QuantScheme::Affine)
.expect("cal");
let mut rng = seeded(7);
for q in params.qmin..=params.qmax {
let on_grid = params.dequantize(q);
let round_trip = params.fake_quant_scalar(on_grid, RoundingMode::Nearest, &mut rng);
assert!(
(round_trip - on_grid).abs() < 1e-9,
"affine grid value {on_grid} (q={q}) -> {round_trip}"
);
}
}
#[test]
fn int8_nearest_error_bounded_by_half_scale() {
let params = QuantParams::from_absmax(2.0, IntDtype::Int8).expect("calibrate");
let half = params.scale / 2.0;
let mut rng = seeded(2);
let mut x = -1.9;
while x <= 1.9 {
let fq = params.fake_quant_scalar(x, RoundingMode::Nearest, &mut rng);
assert!(
(fq - x).abs() <= half + EPS,
"nearest error {} exceeded scale/2 = {half} at x = {x}",
(fq - x).abs()
);
x += 0.013;
}
}
#[test]
fn int4_round_trip_on_grid() {
let params = QuantParams::from_absmax(7.0, IntDtype::Int4).expect("calibrate");
assert!((params.scale - 1.0).abs() < EPS);
let mut rng = seeded(3);
for k in -7..=7 {
let on_grid = k as f64;
let fq = params.fake_quant_scalar(on_grid, RoundingMode::Nearest, &mut rng);
assert!((fq - on_grid).abs() < EPS, "int4 grid {on_grid} -> {fq}");
}
}
#[test]
fn stochastic_rounding_is_unbiased() {
let params = QuantParams::from_absmax(12.7, IntDtype::Int8).expect("calibrate");
let scale = params.scale; let value = 3.0 + 0.7 * scale;
let n: usize = 400_000;
let mut rng = seeded(12345);
let mut sum = 0.0_f64;
let mut saw_lower = false;
let mut saw_upper = false;
let lower = 3.0;
let upper = 3.0 + scale;
for _ in 0..n {
let q = params.fake_quant_scalar(value, RoundingMode::Stochastic, &mut rng);
assert!(
(q - lower).abs() < 1e-9 || (q - upper).abs() < 1e-9,
"stochastic output {q} was not a bracketing grid level"
);
if (q - lower).abs() < 1e-9 {
saw_lower = true;
}
if (q - upper).abs() < 1e-9 {
saw_upper = true;
}
sum += q;
}
let mean = sum / n as f64;
let tolerance = 8.0 * scale / (n as f64).sqrt();
assert!(
(mean - value).abs() < tolerance,
"stochastic mean {mean} deviated from {value} by more than {tolerance}"
);
assert!(saw_lower && saw_upper, "expected both rounding directions");
}
#[test]
fn stochastic_rounding_exact_grid_value_is_stable() {
let params = QuantParams::from_absmax(12.7, IntDtype::Int8).expect("calibrate");
let mut rng = seeded(99);
let exact = 5.0 * params.scale; for _ in 0..1000 {
let q = params.fake_quant_scalar(exact, RoundingMode::Stochastic, &mut rng);
assert!((q - exact).abs() < EPS, "exact grid value drifted: {q}");
}
}
#[test]
fn fp8_constants_match_documentation() {
assert_eq!(E4M3_MAX_NORMAL, 448.0);
assert_eq!(E5M2_MAX_NORMAL, 57344.0);
assert_eq!(Fp8Format::E4M3.max_normal(), 448.0);
assert_eq!(Fp8Format::E5M2.max_normal(), 57344.0);
assert!((Fp8Format::E4M3.min_normal() - 2.0_f64.powi(-6)).abs() < EPS);
assert!((Fp8Format::E4M3.min_subnormal() - 2.0_f64.powi(-9)).abs() < EPS);
assert!((Fp8Format::E5M2.min_normal() - 2.0_f64.powi(-14)).abs() < EPS);
assert!((Fp8Format::E5M2.min_subnormal() - 2.0_f64.powi(-16)).abs() < EPS);
}
#[test]
fn fp8_e4m3_representable_values_map_to_themselves() {
let mut rng = seeded(4);
let representable = [
0.0,
1.0,
1.5, 1.75, 2.0,
-2.0,
0.5,
256.0,
448.0, -448.0,
2.0_f64.powi(-6), 2.0_f64.powi(-9), ];
let input = Array1::from_vec(representable.to_vec());
let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
for (i, (&want, &got)) in representable.iter().zip(out.iter()).enumerate() {
assert!(
(want - got).abs() < EPS,
"E4M3 representable[{i}] = {want} mapped to {got}"
);
}
}
#[test]
fn fp8_e4m3_saturates_to_max_normal() {
let mut rng = seeded(5);
let input = Array1::from_vec(vec![449.0, 1000.0, 1.0e6, -1000.0, f64::INFINITY]);
let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
assert_eq!(out[0], 448.0);
assert_eq!(out[1], 448.0);
assert_eq!(out[2], 448.0);
assert_eq!(out[3], -448.0);
assert_eq!(out[4], 448.0);
}
#[test]
fn fp8_e4m3_nan_propagates() {
let mut rng = seeded(6);
let input = Array1::from_vec(vec![f64::NAN]);
let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
assert!(out[0].is_nan());
}
#[test]
fn fp8_e4m3_rounds_to_nearest_grid_within_half_ulp() {
let mut rng = seeded(8);
let input = Array1::from_vec(vec![1.1]);
let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
assert!((out[0] - 1.125).abs() < EPS, "1.1 -> {}", out[0]);
assert!((out[0] - 1.1).abs() <= 0.125 / 2.0 + EPS);
}
#[test]
fn fp8_e5m2_representable_values_map_to_themselves() {
let mut rng = seeded(9);
let representable = [
0.0,
1.0,
1.5, 2.0,
-4.0,
57344.0, -57344.0,
2.0_f64.powi(-14), 2.0_f64.powi(-16), ];
let input = Array1::from_vec(representable.to_vec());
let out = fake_quant_fp8(&input, Fp8Format::E5M2, RoundingMode::Nearest, &mut rng);
for (i, (&want, &got)) in representable.iter().zip(out.iter()).enumerate() {
assert!(
(want - got).abs() < EPS,
"E5M2 representable[{i}] = {want} mapped to {got}"
);
}
}
#[test]
fn fp8_e5m2_saturates_to_max_normal() {
let mut rng = seeded(10);
let input = Array1::from_vec(vec![60000.0, 1.0e8, -70000.0, f64::INFINITY]);
let out = fake_quant_fp8(&input, Fp8Format::E5M2, RoundingMode::Nearest, &mut rng);
assert_eq!(out[0], 57344.0);
assert_eq!(out[1], 57344.0);
assert_eq!(out[2], -57344.0);
assert_eq!(out[3], 57344.0);
}
#[test]
fn per_channel_scales_differ_and_error_respects_each_scale() {
let x =
Array2::from_shape_vec((2, 3), vec![0.0, 0.5, 1.0, 0.0, 50.0, 100.0]).expect("shape");
let params =
per_channel_params(&x, 0, IntDtype::Int8, QuantScheme::Symmetric).expect("cal");
assert_eq!(params.len(), 2);
assert!(params[1].scale > params[0].scale * 10.0);
assert!((params[0].scale - 1.0 / 127.0).abs() < 1e-6);
assert!((params[1].scale - 100.0 / 127.0).abs() < 1e-6);
let mut rng = seeded(11);
let out = fake_quant_int_per_channel(&x, ¶ms, 0, RoundingMode::Nearest, &mut rng)
.expect("quantize");
for r in 0..2 {
let half = params[r].scale / 2.0;
for c in 0..3 {
let err = (out[[r, c]] - x[[r, c]]).abs();
assert!(
err <= half + EPS,
"per-channel error {err} exceeded scale/2={half} at ({r},{c})"
);
}
}
}
#[test]
fn per_channel_length_mismatch_errors() {
let x = Array2::from_shape_vec((2, 2), vec![1.0, 2.0, 3.0, 4.0]).expect("shape");
let params = per_channel_params(&x, 0, IntDtype::Int8, QuantScheme::Symmetric).expect("c");
let mut rng = seeded(13);
let result =
fake_quant_int_per_channel(&x, ¶ms[..1], 0, RoundingMode::Nearest, &mut rng);
assert!(matches!(
result,
Err(GpuOptimError::DimensionMismatch { .. })
));
}
#[test]
fn ste_passes_gradient_in_range_and_zeros_out_of_range() {
let x = Array1::from_vec(vec![-2.0, -0.5, 0.0, 0.5, 2.0]);
let grad = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, 1.0]);
let out = fake_quant_backward(&grad, &x, -1.0, 1.0).expect("ste");
let expected = [0.0, 1.0, 1.0, 1.0, 0.0];
for (i, (&got, &want)) in out.iter().zip(expected.iter()).enumerate() {
assert!((got - want).abs() < EPS, "STE[{i}] = {got}, want {want}");
}
}
#[test]
fn ste_dimension_mismatch_errors() {
let x = Array1::from_vec(vec![1.0, 2.0, 3.0]);
let grad = Array1::from_vec(vec![1.0, 1.0]);
let result = fake_quant_backward(&grad, &x, -1.0, 1.0);
assert!(matches!(
result,
Err(GpuOptimError::DimensionMismatch { .. })
));
}
#[test]
fn degenerate_calibration_errors() {
let zeros = Array1::from_vec(vec![0.0, 0.0, 0.0]);
assert!(QuantParams::symmetric_from_absmax(&zeros, IntDtype::Int8).is_err());
assert!(
QuantParams::per_tensor_minmax(&zeros, IntDtype::Int8, QuantScheme::Affine).is_err()
);
}
#[test]
fn invalid_bit_width_errors() {
assert!(IntDtype::from_bits(8).is_ok());
assert!(IntDtype::from_bits(4).is_ok());
assert!(matches!(
IntDtype::from_bits(3),
Err(GpuOptimError::UnsupportedOperation(_))
));
assert!(IntDtype::from_bits(16).is_err());
}
#[test]
fn qat_master_update_keeps_fp32_master_and_quantized_view_tracks_it() {
let mut master = Array1::from_vec(vec![0.12, -0.37, 0.88, -0.05, 0.51]);
let master_before = master.clone();
let config = QatConfig::new(
QuantTarget::Int(IntDtype::Int8),
QuantScheme::Symmetric,
RoundingMode::Nearest,
0.1,
);
let mut rng = seeded(2024);
let mut optimizer = QatOptimizer::new(&master, config, &mut rng).expect("new");
let grad = Array1::from_vec(vec![0.1, -0.2, 0.05, 0.3, -0.15]);
optimizer.step(&mut master, &grad, &mut rng).expect("step");
let mut moved = false;
for (&a, &b) in master.iter().zip(master_before.iter()) {
if (a - b).abs() > EPS {
moved = true;
}
}
assert!(moved, "Adam update did not change the master weights");
let params = optimizer.quant_params().expect("int params").to_owned();
let mut check_rng = seeded(2024);
let reference = fake_quant_int(&master, ¶ms, RoundingMode::Nearest, &mut check_rng);
for (&q, &r) in optimizer.quantized_weights().iter().zip(reference.iter()) {
assert!((q - r).abs() < EPS, "quantized view does not track master");
}
let scale = params.scale;
for &q in optimizer.quantized_weights().iter() {
let codes = q / scale;
assert!(
(codes - codes.round()).abs() < 1e-6,
"quantized weight {q} is not on the grid"
);
}
let mut some_off_grid = false;
for &w in master.iter() {
let codes = w / scale;
if (codes - codes.round()).abs() > 1e-6 {
some_off_grid = true;
}
}
assert!(
some_off_grid,
"master weights appear to be quantized (not full precision)"
);
assert_eq!(optimizer.step_count(), 1);
}
#[test]
fn qat_fp8_target_tracks_master() {
let mut master = Array1::from_vec(vec![0.3, -1.2, 4.0, -0.01, 2.5]);
let config = QatConfig::new(
QuantTarget::Fp8(Fp8Format::E4M3),
QuantScheme::Symmetric,
RoundingMode::Nearest,
0.05,
);
let mut rng = seeded(77);
let mut optimizer = QatOptimizer::new(&master, config, &mut rng).expect("new");
assert!(optimizer.quant_params().is_none());
let grad = Array1::from_vec(vec![0.2, 0.1, -0.3, 0.4, -0.05]);
optimizer.step(&mut master, &grad, &mut rng).expect("step");
let mut check_rng = seeded(77);
let reference = fake_quant_fp8(
&master,
Fp8Format::E4M3,
RoundingMode::Nearest,
&mut check_rng,
);
for (&q, &r) in optimizer.quantized_weights().iter().zip(reference.iter()) {
assert!(
(q - r).abs() < EPS,
"fp8 quantized view does not track master"
);
}
}
#[test]
fn percentile_calibration_clips_outliers() {
let mut values = vec![0.0; 100];
for (i, v) in values.iter_mut().enumerate() {
*v = (i as f64) / 100.0; }
values.push(1000.0); let data = Array1::from_vec(values);
let plain = QuantParams::symmetric_from_absmax(&data, IntDtype::Int8).expect("plain");
let clipped =
QuantParams::per_tensor_percentile(&data, IntDtype::Int8, QuantScheme::Symmetric, 0.02)
.expect("clipped");
assert!(
clipped.scale < plain.scale,
"percentile clipping did not reduce the scale ({} vs {})",
clipped.scale,
plain.scale
);
}
}