1use crate::GpuOptimError;
119use scirs2_core::ndarray::{Array1, Array2, Axis, Zip};
120use scirs2_core::random::{Rng, RngExt};
121
122pub const E4M3_MAX_NORMAL: f64 = 448.0;
124
125pub const E5M2_MAX_NORMAL: f64 = 57344.0;
127
128#[derive(Debug, Clone, Copy, PartialEq, Eq)]
131pub enum QuantScheme {
132 Symmetric,
134 Affine,
136}
137
138#[derive(Debug, Clone, Copy, PartialEq, Eq)]
140pub enum IntDtype {
141 Int8,
143 Int4,
145}
146
147impl IntDtype {
148 pub fn bits(self) -> u32 {
150 match self {
151 IntDtype::Int8 => 8,
152 IntDtype::Int4 => 4,
153 }
154 }
155
156 pub fn from_bits(bits: u32) -> Result<Self, GpuOptimError> {
163 match bits {
164 8 => Ok(IntDtype::Int8),
165 4 => Ok(IntDtype::Int4),
166 other => Err(GpuOptimError::UnsupportedOperation(format!(
167 "unsupported integer quantization width: {other} bits (expected 4 or 8)"
168 ))),
169 }
170 }
171
172 pub fn q_range(self, scheme: QuantScheme) -> (i32, i32) {
177 match (self, scheme) {
178 (IntDtype::Int8, QuantScheme::Symmetric) => (-127, 127),
179 (IntDtype::Int8, QuantScheme::Affine) => (-128, 127),
180 (IntDtype::Int4, QuantScheme::Symmetric) => (-7, 7),
181 (IntDtype::Int4, QuantScheme::Affine) => (-8, 7),
182 }
183 }
184}
185
186#[derive(Debug, Clone, Copy, PartialEq, Eq)]
188pub enum Fp8Format {
189 E4M3,
191 E5M2,
193}
194
195impl Fp8Format {
196 pub fn mantissa_bits(self) -> i32 {
198 match self {
199 Fp8Format::E4M3 => 3,
200 Fp8Format::E5M2 => 2,
201 }
202 }
203
204 pub fn exponent_bias(self) -> i32 {
206 match self {
207 Fp8Format::E4M3 => 7,
208 Fp8Format::E5M2 => 15,
209 }
210 }
211
212 pub fn max_normal(self) -> f64 {
214 match self {
215 Fp8Format::E4M3 => E4M3_MAX_NORMAL,
216 Fp8Format::E5M2 => E5M2_MAX_NORMAL,
217 }
218 }
219
220 pub fn min_normal(self) -> f64 {
222 pow2(1 - self.exponent_bias())
223 }
224
225 pub fn min_subnormal(self) -> f64 {
227 pow2(1 - self.exponent_bias() - self.mantissa_bits())
228 }
229}
230
231#[derive(Debug, Clone, Copy, PartialEq, Eq)]
233pub enum RoundingMode {
234 Nearest,
236 Stochastic,
239}
240
241impl RoundingMode {
242 fn round_to_int(self, x: f64, rng: &mut impl Rng) -> f64 {
248 match self {
249 RoundingMode::Nearest => x.round_ties_even(),
250 RoundingMode::Stochastic => {
251 let lower = x.floor();
252 let frac = x - lower;
253 let draw: f64 = rng.random();
254 if draw < frac {
255 lower + 1.0
256 } else {
257 lower
258 }
259 }
260 }
261 }
262}
263
264fn pow2(exp: i32) -> f64 {
269 2.0_f64.powi(exp)
270}
271
272#[derive(Debug, Clone, Copy, PartialEq)]
277pub struct QuantParams {
278 pub scale: f64,
280 pub zero_point: i32,
282 pub qmin: i32,
284 pub qmax: i32,
286}
287
288impl QuantParams {
289 pub fn from_minmax(
301 min_v: f64,
302 max_v: f64,
303 dtype: IntDtype,
304 scheme: QuantScheme,
305 ) -> Result<Self, GpuOptimError> {
306 if !min_v.is_finite() || !max_v.is_finite() || max_v < min_v {
307 return Err(GpuOptimError::InvalidState(format!(
308 "invalid calibration interval: [{min_v}, {max_v}]"
309 )));
310 }
311 match scheme {
312 QuantScheme::Symmetric => Self::from_absmax(min_v.abs().max(max_v.abs()), dtype),
313 QuantScheme::Affine => {
314 let (qmin, qmax) = dtype.q_range(QuantScheme::Affine);
315 let rmin = min_v.min(0.0);
317 let rmax = max_v.max(0.0);
318 let range = rmax - rmin;
319 if range <= 0.0 {
320 return Err(GpuOptimError::InvalidState(
321 "degenerate (zero-range) calibration: real min == max == 0".to_string(),
322 ));
323 }
324 let scale = range / f64::from(qmax - qmin);
325 if !scale.is_finite() || scale <= 0.0 {
326 return Err(GpuOptimError::InvalidState(format!(
327 "degenerate affine scale derived from interval [{min_v}, {max_v}]"
328 )));
329 }
330 let zp = (f64::from(qmin) - rmin / scale)
331 .round_ties_even()
332 .clamp(f64::from(qmin), f64::from(qmax));
333 Ok(Self {
334 scale,
335 zero_point: zp as i32,
336 qmin,
337 qmax,
338 })
339 }
340 }
341 }
342
343 pub fn from_absmax(absmax: f64, dtype: IntDtype) -> Result<Self, GpuOptimError> {
352 let (qmin, qmax) = dtype.q_range(QuantScheme::Symmetric);
353 if !absmax.is_finite() || absmax <= 0.0 {
354 return Err(GpuOptimError::InvalidState(
355 "degenerate (zero-range) calibration: absmax must be finite and > 0".to_string(),
356 ));
357 }
358 let scale = absmax / f64::from(qmax);
359 if !scale.is_finite() || scale <= 0.0 {
360 return Err(GpuOptimError::InvalidState(
361 "degenerate symmetric scale".to_string(),
362 ));
363 }
364 Ok(Self {
365 scale,
366 zero_point: 0,
367 qmin,
368 qmax,
369 })
370 }
371
372 pub fn per_tensor_minmax(
379 x: &Array1<f64>,
380 dtype: IntDtype,
381 scheme: QuantScheme,
382 ) -> Result<Self, GpuOptimError> {
383 let (min_v, max_v) = finite_min_max(x.iter().copied())?;
384 Self::from_minmax(min_v, max_v, dtype, scheme)
385 }
386
387 pub fn symmetric_from_absmax(x: &Array1<f64>, dtype: IntDtype) -> Result<Self, GpuOptimError> {
393 if x.is_empty() {
394 return Err(GpuOptimError::InvalidState(
395 "cannot calibrate an empty tensor".to_string(),
396 ));
397 }
398 let mut absmax = 0.0_f64;
399 for &v in x.iter() {
400 if !v.is_finite() {
401 return Err(GpuOptimError::InvalidState(
402 "non-finite value in calibration tensor".to_string(),
403 ));
404 }
405 absmax = absmax.max(v.abs());
406 }
407 Self::from_absmax(absmax, dtype)
408 }
409
410 pub fn per_tensor_percentile(
421 x: &Array1<f64>,
422 dtype: IntDtype,
423 scheme: QuantScheme,
424 clip_fraction: f64,
425 ) -> Result<Self, GpuOptimError> {
426 if !(0.0..0.5).contains(&clip_fraction) {
427 return Err(GpuOptimError::InvalidState(format!(
428 "clip_fraction must lie in [0, 0.5), got {clip_fraction}"
429 )));
430 }
431 if x.is_empty() {
432 return Err(GpuOptimError::InvalidState(
433 "cannot calibrate an empty tensor".to_string(),
434 ));
435 }
436 let mut sorted: Vec<f64> = Vec::with_capacity(x.len());
437 for &v in x.iter() {
438 if !v.is_finite() {
439 return Err(GpuOptimError::InvalidState(
440 "non-finite value in calibration tensor".to_string(),
441 ));
442 }
443 sorted.push(v);
444 }
445 sorted.sort_by(|a, b| a.total_cmp(b));
446 let last = sorted.len() - 1;
447 let lo_idx = (clip_fraction * last as f64).floor() as usize;
448 let hi_idx = ((1.0 - clip_fraction) * last as f64).ceil() as usize;
449 let min_v = sorted[lo_idx.min(last)];
450 let max_v = sorted[hi_idx.min(last)];
451 Self::from_minmax(min_v, max_v, dtype, scheme)
452 }
453
454 pub fn quantize(&self, x: f64, mode: RoundingMode, rng: &mut impl Rng) -> i32 {
456 let scaled = x / self.scale + f64::from(self.zero_point);
457 let rounded = mode
458 .round_to_int(scaled, rng)
459 .clamp(f64::from(self.qmin), f64::from(self.qmax));
460 rounded as i32
461 }
462
463 pub fn dequantize(&self, q: i32) -> f64 {
465 f64::from(q - self.zero_point) * self.scale
466 }
467
468 pub fn fake_quant_scalar(&self, x: f64, mode: RoundingMode, rng: &mut impl Rng) -> f64 {
470 self.dequantize(self.quantize(x, mode, rng))
471 }
472
473 pub fn real_min(&self) -> f64 {
476 self.dequantize(self.qmin)
477 }
478
479 pub fn real_max(&self) -> f64 {
481 self.dequantize(self.qmax)
482 }
483}
484
485fn finite_min_max(values: impl Iterator<Item = f64>) -> Result<(f64, f64), GpuOptimError> {
487 let mut min_v = f64::INFINITY;
488 let mut max_v = f64::NEG_INFINITY;
489 let mut count = 0_usize;
490 for v in values {
491 if !v.is_finite() {
492 return Err(GpuOptimError::InvalidState(
493 "non-finite value in calibration tensor".to_string(),
494 ));
495 }
496 min_v = min_v.min(v);
497 max_v = max_v.max(v);
498 count += 1;
499 }
500 if count == 0 {
501 return Err(GpuOptimError::InvalidState(
502 "cannot calibrate an empty tensor".to_string(),
503 ));
504 }
505 Ok((min_v, max_v))
506}
507
508pub fn fake_quant_int(
513 x: &Array1<f64>,
514 params: &QuantParams,
515 mode: RoundingMode,
516 rng: &mut impl Rng,
517) -> Array1<f64> {
518 let mut out = Vec::with_capacity(x.len());
519 for &v in x.iter() {
520 out.push(params.fake_quant_scalar(v, mode, rng));
521 }
522 Array1::from_vec(out)
523}
524
525pub fn per_channel_params(
536 x: &Array2<f64>,
537 axis: usize,
538 dtype: IntDtype,
539 scheme: QuantScheme,
540) -> Result<Vec<QuantParams>, GpuOptimError> {
541 if axis > 1 {
542 return Err(GpuOptimError::InvalidState(format!(
543 "per-channel axis must be 0 or 1, got {axis}"
544 )));
545 }
546 let n_channels = x.shape()[axis];
547 let mut params = Vec::with_capacity(n_channels);
548 for channel in 0..n_channels {
549 let lane = x.index_axis(Axis(axis), channel);
550 let (min_v, max_v) = finite_min_max(lane.iter().copied())?;
551 params.push(QuantParams::from_minmax(min_v, max_v, dtype, scheme)?);
552 }
553 Ok(params)
554}
555
556pub fn fake_quant_int_per_channel(
567 x: &Array2<f64>,
568 params: &[QuantParams],
569 axis: usize,
570 mode: RoundingMode,
571 rng: &mut impl Rng,
572) -> Result<Array2<f64>, GpuOptimError> {
573 if axis > 1 {
574 return Err(GpuOptimError::InvalidState(format!(
575 "per-channel axis must be 0 or 1, got {axis}"
576 )));
577 }
578 let n_channels = x.shape()[axis];
579 if params.len() != n_channels {
580 return Err(GpuOptimError::DimensionMismatch {
581 expected: vec![n_channels],
582 actual: vec![params.len()],
583 });
584 }
585 let (n_rows, n_cols) = (x.shape()[0], x.shape()[1]);
586 let mut out = Array2::<f64>::zeros((n_rows, n_cols));
587 for r in 0..n_rows {
588 for c in 0..n_cols {
589 let channel = if axis == 0 { r } else { c };
590 let p = ¶ms[channel];
591 out[[r, c]] = p.fake_quant_scalar(x[[r, c]], mode, rng);
592 }
593 }
594 Ok(out)
595}
596
597fn fp8_quantize_magnitude(
602 a: f64,
603 mantissa_bits: i32,
604 exponent_bias: i32,
605 max_normal: f64,
606 mode: RoundingMode,
607 rng: &mut impl Rng,
608) -> f64 {
609 if a == 0.0 {
610 return 0.0;
611 }
612 let emin = 1 - exponent_bias;
614 let e = ((a.to_bits() >> 52) & 0x7ff) as i32 - 1023;
619 let step_exp = e.max(emin) - mantissa_bits;
620 let step = pow2(step_exp);
621 let ratio = a / step;
622 let rounded = mode.round_to_int(ratio, rng);
623 let magnitude = rounded * step;
624 if magnitude > max_normal {
625 max_normal
626 } else {
627 magnitude
628 }
629}
630
631pub fn fake_quant_fp8(
637 x: &Array1<f64>,
638 format: Fp8Format,
639 mode: RoundingMode,
640 rng: &mut impl Rng,
641) -> Array1<f64> {
642 let mantissa_bits = format.mantissa_bits();
643 let exponent_bias = format.exponent_bias();
644 let max_normal = format.max_normal();
645 let mut out = Vec::with_capacity(x.len());
646 for &v in x.iter() {
647 let q = if v.is_nan() {
648 f64::NAN
649 } else if v.is_infinite() {
650 max_normal.copysign(v)
651 } else {
652 let magnitude = fp8_quantize_magnitude(
653 v.abs(),
654 mantissa_bits,
655 exponent_bias,
656 max_normal,
657 mode,
658 rng,
659 );
660 magnitude.copysign(v)
661 };
662 out.push(q);
663 }
664 Array1::from_vec(out)
665}
666
667pub fn fake_quant_backward(
678 grad: &Array1<f64>,
679 x: &Array1<f64>,
680 qmin_real: f64,
681 qmax_real: f64,
682) -> Result<Array1<f64>, GpuOptimError> {
683 if grad.len() != x.len() {
684 return Err(GpuOptimError::DimensionMismatch {
685 expected: vec![x.len()],
686 actual: vec![grad.len()],
687 });
688 }
689 let (lo, hi) = if qmin_real <= qmax_real {
690 (qmin_real, qmax_real)
691 } else {
692 (qmax_real, qmin_real)
693 };
694 let mut out = Vec::with_capacity(grad.len());
695 for (&g, &v) in grad.iter().zip(x.iter()) {
696 if v >= lo && v <= hi {
697 out.push(g);
698 } else {
699 out.push(0.0);
700 }
701 }
702 Ok(Array1::from_vec(out))
703}
704
705#[derive(Debug, Clone, Copy, PartialEq, Eq)]
707pub enum QuantTarget {
708 Int(IntDtype),
710 Fp8(Fp8Format),
712}
713
714#[derive(Debug, Clone, Copy, PartialEq)]
717pub struct QatConfig {
718 pub target: QuantTarget,
720 pub scheme: QuantScheme,
722 pub rounding: RoundingMode,
724 pub lr: f64,
726 pub beta1: f64,
728 pub beta2: f64,
730 pub eps: f64,
732 pub weight_decay: f64,
734}
735
736impl QatConfig {
737 pub fn new(target: QuantTarget, scheme: QuantScheme, rounding: RoundingMode, lr: f64) -> Self {
740 Self {
741 target,
742 scheme,
743 rounding,
744 lr,
745 beta1: 0.9,
746 beta2: 0.999,
747 eps: 1e-8,
748 weight_decay: 0.0,
749 }
750 }
751}
752
753#[derive(Debug, Clone)]
760pub struct QatOptimizer {
761 config: QatConfig,
762 first_moment: Array1<f64>,
763 second_moment: Array1<f64>,
764 step_count: u64,
765 quantized: Array1<f64>,
766 int_params: Option<QuantParams>,
767}
768
769impl QatOptimizer {
770 pub fn new(
779 master: &Array1<f64>,
780 config: QatConfig,
781 rng: &mut impl Rng,
782 ) -> Result<Self, GpuOptimError> {
783 if master.is_empty() {
784 return Err(GpuOptimError::InvalidState(
785 "cannot construct a QatOptimizer over empty master weights".to_string(),
786 ));
787 }
788 let n = master.len();
789 let mut optimizer = Self {
790 config,
791 first_moment: Array1::zeros(n),
792 second_moment: Array1::zeros(n),
793 step_count: 0,
794 quantized: Array1::zeros(n),
795 int_params: None,
796 };
797 optimizer.requantize(master, rng)?;
798 Ok(optimizer)
799 }
800
801 pub fn quantized_weights(&self) -> &Array1<f64> {
804 &self.quantized
805 }
806
807 pub fn quant_params(&self) -> Option<&QuantParams> {
810 self.int_params.as_ref()
811 }
812
813 pub fn step_count(&self) -> u64 {
815 self.step_count
816 }
817
818 pub fn step(
830 &mut self,
831 master: &mut Array1<f64>,
832 grad: &Array1<f64>,
833 rng: &mut impl Rng,
834 ) -> Result<(), GpuOptimError> {
835 if master.len() != grad.len() {
836 return Err(GpuOptimError::DimensionMismatch {
837 expected: vec![master.len()],
838 actual: vec![grad.len()],
839 });
840 }
841 if master.len() != self.first_moment.len() {
842 return Err(GpuOptimError::DimensionMismatch {
843 expected: vec![self.first_moment.len()],
844 actual: vec![master.len()],
845 });
846 }
847
848 self.step_count += 1;
849 let t = self.step_count as i32;
850 let beta1 = self.config.beta1;
851 let beta2 = self.config.beta2;
852 let lr = self.config.lr;
853 let eps = self.config.eps;
854 let weight_decay = self.config.weight_decay;
855 let bias_correction1 = 1.0 - beta1.powi(t);
856 let bias_correction2 = 1.0 - beta2.powi(t);
857
858 Zip::from(&mut *master)
859 .and(grad)
860 .and(&mut self.first_moment)
861 .and(&mut self.second_moment)
862 .for_each(|weight, &g, m, v| {
863 *m = beta1 * *m + (1.0 - beta1) * g;
864 *v = beta2 * *v + (1.0 - beta2) * g * g;
865 let m_hat = *m / bias_correction1;
866 let v_hat = *v / bias_correction2;
867 if weight_decay != 0.0 {
869 *weight -= lr * weight_decay * *weight;
870 }
871 *weight -= lr * m_hat / (v_hat.sqrt() + eps);
872 });
873
874 self.requantize(master, rng)?;
875 Ok(())
876 }
877
878 fn requantize(
880 &mut self,
881 master: &Array1<f64>,
882 rng: &mut impl Rng,
883 ) -> Result<(), GpuOptimError> {
884 match self.config.target {
885 QuantTarget::Int(dtype) => {
886 let params = QuantParams::per_tensor_minmax(master, dtype, self.config.scheme)?;
887 self.quantized = fake_quant_int(master, ¶ms, self.config.rounding, rng);
888 self.int_params = Some(params);
889 }
890 QuantTarget::Fp8(format) => {
891 self.quantized = fake_quant_fp8(master, format, self.config.rounding, rng);
892 self.int_params = None;
893 }
894 }
895 Ok(())
896 }
897}
898
899#[cfg(test)]
900mod tests {
901 use super::*;
902 use scirs2_core::random::Random;
903
904 const EPS: f64 = 1e-9;
905
906 fn seeded(seed: u64) -> Random<scirs2_core::random::rngs::StdRng> {
907 Random::seed(seed)
908 }
909
910 #[test]
911 fn int8_grid_values_quantize_to_themselves() {
912 let params = QuantParams::from_absmax(12.7, IntDtype::Int8).expect("calibrate");
914 assert!((params.scale - 0.1).abs() < EPS);
915 assert_eq!(params.zero_point, 0);
916 let mut rng = seeded(1);
917 for k in -120..=120 {
919 let on_grid = k as f64 * params.scale;
920 let round_trip = params.fake_quant_scalar(on_grid, RoundingMode::Nearest, &mut rng);
921 assert!(
922 (round_trip - on_grid).abs() < EPS,
923 "grid value {on_grid} did not map to itself (got {round_trip})"
924 );
925 }
926 }
927
928 #[test]
929 fn int8_affine_grid_values_quantize_to_themselves() {
930 let data = Array1::from_vec(vec![-0.3, 1.7, 0.0, 0.9, -0.1]);
931 let params = QuantParams::per_tensor_minmax(&data, IntDtype::Int8, QuantScheme::Affine)
932 .expect("cal");
933 let mut rng = seeded(7);
934 for q in params.qmin..=params.qmax {
936 let on_grid = params.dequantize(q);
937 let round_trip = params.fake_quant_scalar(on_grid, RoundingMode::Nearest, &mut rng);
938 assert!(
939 (round_trip - on_grid).abs() < 1e-9,
940 "affine grid value {on_grid} (q={q}) -> {round_trip}"
941 );
942 }
943 }
944
945 #[test]
946 fn int8_nearest_error_bounded_by_half_scale() {
947 let params = QuantParams::from_absmax(2.0, IntDtype::Int8).expect("calibrate");
948 let half = params.scale / 2.0;
949 let mut rng = seeded(2);
950 let mut x = -1.9;
952 while x <= 1.9 {
953 let fq = params.fake_quant_scalar(x, RoundingMode::Nearest, &mut rng);
954 assert!(
955 (fq - x).abs() <= half + EPS,
956 "nearest error {} exceeded scale/2 = {half} at x = {x}",
957 (fq - x).abs()
958 );
959 x += 0.013;
960 }
961 }
962
963 #[test]
964 fn int4_round_trip_on_grid() {
965 let params = QuantParams::from_absmax(7.0, IntDtype::Int4).expect("calibrate");
966 assert!((params.scale - 1.0).abs() < EPS);
968 let mut rng = seeded(3);
969 for k in -7..=7 {
970 let on_grid = k as f64;
971 let fq = params.fake_quant_scalar(on_grid, RoundingMode::Nearest, &mut rng);
972 assert!((fq - on_grid).abs() < EPS, "int4 grid {on_grid} -> {fq}");
973 }
974 }
975
976 #[test]
977 fn stochastic_rounding_is_unbiased() {
978 let params = QuantParams::from_absmax(12.7, IntDtype::Int8).expect("calibrate");
979 let scale = params.scale; let value = 3.0 + 0.7 * scale;
982 let n: usize = 400_000;
983 let mut rng = seeded(12345);
984 let mut sum = 0.0_f64;
985 let mut saw_lower = false;
986 let mut saw_upper = false;
987 let lower = 3.0;
988 let upper = 3.0 + scale;
989 for _ in 0..n {
990 let q = params.fake_quant_scalar(value, RoundingMode::Stochastic, &mut rng);
991 assert!(
993 (q - lower).abs() < 1e-9 || (q - upper).abs() < 1e-9,
994 "stochastic output {q} was not a bracketing grid level"
995 );
996 if (q - lower).abs() < 1e-9 {
997 saw_lower = true;
998 }
999 if (q - upper).abs() < 1e-9 {
1000 saw_upper = true;
1001 }
1002 sum += q;
1003 }
1004 let mean = sum / n as f64;
1005 let tolerance = 8.0 * scale / (n as f64).sqrt();
1007 assert!(
1008 (mean - value).abs() < tolerance,
1009 "stochastic mean {mean} deviated from {value} by more than {tolerance}"
1010 );
1011 assert!(saw_lower && saw_upper, "expected both rounding directions");
1012 }
1013
1014 #[test]
1015 fn stochastic_rounding_exact_grid_value_is_stable() {
1016 let params = QuantParams::from_absmax(12.7, IntDtype::Int8).expect("calibrate");
1017 let mut rng = seeded(99);
1018 let exact = 5.0 * params.scale; for _ in 0..1000 {
1020 let q = params.fake_quant_scalar(exact, RoundingMode::Stochastic, &mut rng);
1021 assert!((q - exact).abs() < EPS, "exact grid value drifted: {q}");
1022 }
1023 }
1024
1025 #[test]
1026 fn fp8_constants_match_documentation() {
1027 assert_eq!(E4M3_MAX_NORMAL, 448.0);
1028 assert_eq!(E5M2_MAX_NORMAL, 57344.0);
1029 assert_eq!(Fp8Format::E4M3.max_normal(), 448.0);
1030 assert_eq!(Fp8Format::E5M2.max_normal(), 57344.0);
1031 assert!((Fp8Format::E4M3.min_normal() - 2.0_f64.powi(-6)).abs() < EPS);
1032 assert!((Fp8Format::E4M3.min_subnormal() - 2.0_f64.powi(-9)).abs() < EPS);
1033 assert!((Fp8Format::E5M2.min_normal() - 2.0_f64.powi(-14)).abs() < EPS);
1034 assert!((Fp8Format::E5M2.min_subnormal() - 2.0_f64.powi(-16)).abs() < EPS);
1035 }
1036
1037 #[test]
1038 fn fp8_e4m3_representable_values_map_to_themselves() {
1039 let mut rng = seeded(4);
1040 let representable = [
1041 0.0,
1042 1.0,
1043 1.5, 1.75, 2.0,
1046 -2.0,
1047 0.5,
1048 256.0,
1049 448.0, -448.0,
1051 2.0_f64.powi(-6), 2.0_f64.powi(-9), ];
1054 let input = Array1::from_vec(representable.to_vec());
1055 let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
1056 for (i, (&want, &got)) in representable.iter().zip(out.iter()).enumerate() {
1057 assert!(
1058 (want - got).abs() < EPS,
1059 "E4M3 representable[{i}] = {want} mapped to {got}"
1060 );
1061 }
1062 }
1063
1064 #[test]
1065 fn fp8_e4m3_saturates_to_max_normal() {
1066 let mut rng = seeded(5);
1067 let input = Array1::from_vec(vec![449.0, 1000.0, 1.0e6, -1000.0, f64::INFINITY]);
1068 let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
1069 assert_eq!(out[0], 448.0);
1070 assert_eq!(out[1], 448.0);
1071 assert_eq!(out[2], 448.0);
1072 assert_eq!(out[3], -448.0);
1073 assert_eq!(out[4], 448.0);
1074 }
1075
1076 #[test]
1077 fn fp8_e4m3_nan_propagates() {
1078 let mut rng = seeded(6);
1079 let input = Array1::from_vec(vec![f64::NAN]);
1080 let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
1081 assert!(out[0].is_nan());
1082 }
1083
1084 #[test]
1085 fn fp8_e4m3_rounds_to_nearest_grid_within_half_ulp() {
1086 let mut rng = seeded(8);
1087 let input = Array1::from_vec(vec![1.1]);
1089 let out = fake_quant_fp8(&input, Fp8Format::E4M3, RoundingMode::Nearest, &mut rng);
1090 assert!((out[0] - 1.125).abs() < EPS, "1.1 -> {}", out[0]);
1091 assert!((out[0] - 1.1).abs() <= 0.125 / 2.0 + EPS);
1092 }
1093
1094 #[test]
1095 fn fp8_e5m2_representable_values_map_to_themselves() {
1096 let mut rng = seeded(9);
1097 let representable = [
1098 0.0,
1099 1.0,
1100 1.5, 2.0,
1102 -4.0,
1103 57344.0, -57344.0,
1105 2.0_f64.powi(-14), 2.0_f64.powi(-16), ];
1108 let input = Array1::from_vec(representable.to_vec());
1109 let out = fake_quant_fp8(&input, Fp8Format::E5M2, RoundingMode::Nearest, &mut rng);
1110 for (i, (&want, &got)) in representable.iter().zip(out.iter()).enumerate() {
1111 assert!(
1112 (want - got).abs() < EPS,
1113 "E5M2 representable[{i}] = {want} mapped to {got}"
1114 );
1115 }
1116 }
1117
1118 #[test]
1119 fn fp8_e5m2_saturates_to_max_normal() {
1120 let mut rng = seeded(10);
1121 let input = Array1::from_vec(vec![60000.0, 1.0e8, -70000.0, f64::INFINITY]);
1122 let out = fake_quant_fp8(&input, Fp8Format::E5M2, RoundingMode::Nearest, &mut rng);
1123 assert_eq!(out[0], 57344.0);
1124 assert_eq!(out[1], 57344.0);
1125 assert_eq!(out[2], -57344.0);
1126 assert_eq!(out[3], 57344.0);
1127 }
1128
1129 #[test]
1130 fn per_channel_scales_differ_and_error_respects_each_scale() {
1131 let x =
1133 Array2::from_shape_vec((2, 3), vec![0.0, 0.5, 1.0, 0.0, 50.0, 100.0]).expect("shape");
1134 let params =
1135 per_channel_params(&x, 0, IntDtype::Int8, QuantScheme::Symmetric).expect("cal");
1136 assert_eq!(params.len(), 2);
1137 assert!(params[1].scale > params[0].scale * 10.0);
1139 assert!((params[0].scale - 1.0 / 127.0).abs() < 1e-6);
1140 assert!((params[1].scale - 100.0 / 127.0).abs() < 1e-6);
1141
1142 let mut rng = seeded(11);
1143 let out = fake_quant_int_per_channel(&x, ¶ms, 0, RoundingMode::Nearest, &mut rng)
1144 .expect("quantize");
1145 for r in 0..2 {
1146 let half = params[r].scale / 2.0;
1147 for c in 0..3 {
1148 let err = (out[[r, c]] - x[[r, c]]).abs();
1149 assert!(
1150 err <= half + EPS,
1151 "per-channel error {err} exceeded scale/2={half} at ({r},{c})"
1152 );
1153 }
1154 }
1155 }
1156
1157 #[test]
1158 fn per_channel_length_mismatch_errors() {
1159 let x = Array2::from_shape_vec((2, 2), vec![1.0, 2.0, 3.0, 4.0]).expect("shape");
1160 let params = per_channel_params(&x, 0, IntDtype::Int8, QuantScheme::Symmetric).expect("c");
1161 let mut rng = seeded(13);
1162 let result =
1164 fake_quant_int_per_channel(&x, ¶ms[..1], 0, RoundingMode::Nearest, &mut rng);
1165 assert!(matches!(
1166 result,
1167 Err(GpuOptimError::DimensionMismatch { .. })
1168 ));
1169 }
1170
1171 #[test]
1172 fn ste_passes_gradient_in_range_and_zeros_out_of_range() {
1173 let x = Array1::from_vec(vec![-2.0, -0.5, 0.0, 0.5, 2.0]);
1174 let grad = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0, 1.0]);
1175 let out = fake_quant_backward(&grad, &x, -1.0, 1.0).expect("ste");
1176 let expected = [0.0, 1.0, 1.0, 1.0, 0.0];
1177 for (i, (&got, &want)) in out.iter().zip(expected.iter()).enumerate() {
1178 assert!((got - want).abs() < EPS, "STE[{i}] = {got}, want {want}");
1179 }
1180 }
1181
1182 #[test]
1183 fn ste_dimension_mismatch_errors() {
1184 let x = Array1::from_vec(vec![1.0, 2.0, 3.0]);
1185 let grad = Array1::from_vec(vec![1.0, 1.0]);
1186 let result = fake_quant_backward(&grad, &x, -1.0, 1.0);
1187 assert!(matches!(
1188 result,
1189 Err(GpuOptimError::DimensionMismatch { .. })
1190 ));
1191 }
1192
1193 #[test]
1194 fn degenerate_calibration_errors() {
1195 let zeros = Array1::from_vec(vec![0.0, 0.0, 0.0]);
1196 assert!(QuantParams::symmetric_from_absmax(&zeros, IntDtype::Int8).is_err());
1197 assert!(
1198 QuantParams::per_tensor_minmax(&zeros, IntDtype::Int8, QuantScheme::Affine).is_err()
1199 );
1200 }
1201
1202 #[test]
1203 fn invalid_bit_width_errors() {
1204 assert!(IntDtype::from_bits(8).is_ok());
1205 assert!(IntDtype::from_bits(4).is_ok());
1206 assert!(matches!(
1207 IntDtype::from_bits(3),
1208 Err(GpuOptimError::UnsupportedOperation(_))
1209 ));
1210 assert!(IntDtype::from_bits(16).is_err());
1211 }
1212
1213 #[test]
1214 fn qat_master_update_keeps_fp32_master_and_quantized_view_tracks_it() {
1215 let mut master = Array1::from_vec(vec![0.12, -0.37, 0.88, -0.05, 0.51]);
1216 let master_before = master.clone();
1217 let config = QatConfig::new(
1218 QuantTarget::Int(IntDtype::Int8),
1219 QuantScheme::Symmetric,
1220 RoundingMode::Nearest,
1221 0.1,
1222 );
1223 let mut rng = seeded(2024);
1224 let mut optimizer = QatOptimizer::new(&master, config, &mut rng).expect("new");
1225
1226 let grad = Array1::from_vec(vec![0.1, -0.2, 0.05, 0.3, -0.15]);
1227 optimizer.step(&mut master, &grad, &mut rng).expect("step");
1228
1229 let mut moved = false;
1231 for (&a, &b) in master.iter().zip(master_before.iter()) {
1232 if (a - b).abs() > EPS {
1233 moved = true;
1234 }
1235 }
1236 assert!(moved, "Adam update did not change the master weights");
1237
1238 let params = optimizer.quant_params().expect("int params").to_owned();
1240 let mut check_rng = seeded(2024);
1241 let reference = fake_quant_int(&master, ¶ms, RoundingMode::Nearest, &mut check_rng);
1242 for (&q, &r) in optimizer.quantized_weights().iter().zip(reference.iter()) {
1243 assert!((q - r).abs() < EPS, "quantized view does not track master");
1244 }
1245
1246 let scale = params.scale;
1250 for &q in optimizer.quantized_weights().iter() {
1251 let codes = q / scale;
1252 assert!(
1253 (codes - codes.round()).abs() < 1e-6,
1254 "quantized weight {q} is not on the grid"
1255 );
1256 }
1257 let mut some_off_grid = false;
1258 for &w in master.iter() {
1259 let codes = w / scale;
1260 if (codes - codes.round()).abs() > 1e-6 {
1261 some_off_grid = true;
1262 }
1263 }
1264 assert!(
1265 some_off_grid,
1266 "master weights appear to be quantized (not full precision)"
1267 );
1268
1269 assert_eq!(optimizer.step_count(), 1);
1270 }
1271
1272 #[test]
1273 fn qat_fp8_target_tracks_master() {
1274 let mut master = Array1::from_vec(vec![0.3, -1.2, 4.0, -0.01, 2.5]);
1275 let config = QatConfig::new(
1276 QuantTarget::Fp8(Fp8Format::E4M3),
1277 QuantScheme::Symmetric,
1278 RoundingMode::Nearest,
1279 0.05,
1280 );
1281 let mut rng = seeded(77);
1282 let mut optimizer = QatOptimizer::new(&master, config, &mut rng).expect("new");
1283 assert!(optimizer.quant_params().is_none());
1284
1285 let grad = Array1::from_vec(vec![0.2, 0.1, -0.3, 0.4, -0.05]);
1286 optimizer.step(&mut master, &grad, &mut rng).expect("step");
1287
1288 let mut check_rng = seeded(77);
1289 let reference = fake_quant_fp8(
1291 &master,
1292 Fp8Format::E4M3,
1293 RoundingMode::Nearest,
1294 &mut check_rng,
1295 );
1296 for (&q, &r) in optimizer.quantized_weights().iter().zip(reference.iter()) {
1297 assert!(
1298 (q - r).abs() < EPS,
1299 "fp8 quantized view does not track master"
1300 );
1301 }
1302 }
1303
1304 #[test]
1305 fn percentile_calibration_clips_outliers() {
1306 let mut values = vec![0.0; 100];
1309 for (i, v) in values.iter_mut().enumerate() {
1310 *v = (i as f64) / 100.0; }
1312 values.push(1000.0); let data = Array1::from_vec(values);
1314 let plain = QuantParams::symmetric_from_absmax(&data, IntDtype::Int8).expect("plain");
1315 let clipped =
1316 QuantParams::per_tensor_percentile(&data, IntDtype::Int8, QuantScheme::Symmetric, 0.02)
1317 .expect("clipped");
1318 assert!(
1319 clipped.scale < plain.scale,
1320 "percentile clipping did not reduce the scale ({} vs {})",
1321 clipped.scale,
1322 plain.scale
1323 );
1324 }
1325}