1use std::collections::HashMap;
24
25use crate::dtype::Float;
26use crate::error::{FerrotorchError, FerrotorchResult};
27use crate::storage::TensorStorage;
28use crate::tensor::Tensor;
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
36pub enum QuantScheme {
37 PerTensor,
39 PerChannel(usize),
41}
42
43#[derive(Debug, Clone, Copy, PartialEq, Eq)]
45pub enum QuantDtype {
46 Int8,
48 Int4,
50 Uint8,
52}
53
54impl QuantDtype {
55 #[inline]
57 fn qmin(self) -> i32 {
58 match self {
59 QuantDtype::Int8 => -128,
60 QuantDtype::Int4 => -8,
61 QuantDtype::Uint8 => 0,
62 }
63 }
64
65 #[inline]
67 fn qmax(self) -> i32 {
68 match self {
69 QuantDtype::Int8 => 127,
70 QuantDtype::Int4 => 7,
71 QuantDtype::Uint8 => 255,
72 }
73 }
74}
75
76#[derive(Debug, Clone)]
88pub struct QuantizedTensor {
89 data: Vec<i8>,
93 scale: Vec<f32>,
95 zero_point: Vec<i32>,
97 shape: Vec<usize>,
99 scheme: QuantScheme,
101 dtype: QuantDtype,
103}
104
105impl QuantizedTensor {
106 #[inline]
108 pub fn numel(&self) -> usize {
109 self.shape.iter().product()
110 }
111
112 #[inline]
114 pub fn shape(&self) -> &[usize] {
115 &self.shape
116 }
117
118 #[inline]
120 pub fn data(&self) -> &[i8] {
121 &self.data
122 }
123
124 #[inline]
126 pub fn scale(&self) -> &[f32] {
127 &self.scale
128 }
129
130 #[inline]
132 pub fn zero_point(&self) -> &[i32] {
133 &self.zero_point
134 }
135
136 #[inline]
138 pub fn scheme(&self) -> QuantScheme {
139 self.scheme
140 }
141
142 #[inline]
144 pub fn qdtype(&self) -> QuantDtype {
145 self.dtype
146 }
147}
148
149fn compute_scale_zp(min_val: f32, max_val: f32, dtype: QuantDtype) -> (f32, i32) {
164 let qmin = dtype.qmin();
165 let qmax = dtype.qmax();
166
167 let min_val = min_val.min(0.0);
169 let max_val = max_val.max(0.0);
170
171 let range = (max_val - min_val).max(f32::EPSILON);
174 let scale = range / (qmax - qmin) as f32;
175
176 let zp = (qmin as f32 - min_val / scale).round() as i32;
181
182 (scale, zp)
183}
184
185#[inline]
191fn quantize_val(x: f32, scale: f32, zp: i32, qmin: i32, qmax: i32, is_unsigned: bool) -> i8 {
192 let q = (x / scale + zp as f32).round() as i32;
193 let clamped = q.clamp(qmin, qmax);
194 if is_unsigned {
195 (clamped as u8) as i8
196 } else {
197 clamped as i8
198 }
199}
200
201#[inline]
204fn stored_to_i32(val: i8, is_unsigned: bool) -> i32 {
205 if is_unsigned {
206 (val as u8) as i32
207 } else {
208 val as i32
209 }
210}
211
212#[inline]
217fn channel_index(flat_index: usize, shape: &[usize], axis: usize) -> usize {
218 let stride: usize = shape[axis + 1..].iter().product();
220 (flat_index / stride) % shape[axis]
221}
222
223pub fn quantize<T: Float>(
238 tensor: &Tensor<T>,
239 scheme: QuantScheme,
240 dtype: QuantDtype,
241) -> FerrotorchResult<QuantizedTensor> {
242 let data = tensor.data()?;
243 let shape = tensor.shape().to_vec();
244 let numel = tensor.numel();
245 let qmin = dtype.qmin();
246 let qmax = dtype.qmax();
247
248 let is_unsigned = dtype == QuantDtype::Uint8;
249
250 match scheme {
251 QuantScheme::PerTensor => {
252 let mut min_val = f32::INFINITY;
254 let mut max_val = f32::NEG_INFINITY;
255 for &v in data {
256 let f = v.to_f32().unwrap();
257 if f < min_val {
258 min_val = f;
259 }
260 if f > max_val {
261 max_val = f;
262 }
263 }
264
265 let (scale, zp) = compute_scale_zp(min_val, max_val, dtype);
266
267 let qdata: Vec<i8> = data
268 .iter()
269 .map(|&v| quantize_val(v.to_f32().unwrap(), scale, zp, qmin, qmax, is_unsigned))
270 .collect();
271
272 Ok(QuantizedTensor {
273 data: qdata,
274 scale: vec![scale],
275 zero_point: vec![zp],
276 shape,
277 scheme,
278 dtype,
279 })
280 }
281
282 QuantScheme::PerChannel(axis) => {
283 if axis >= shape.len() {
284 return Err(FerrotorchError::InvalidArgument {
285 message: format!(
286 "PerChannel axis {axis} out of range for {}-d tensor",
287 shape.len()
288 ),
289 });
290 }
291
292 let num_channels = shape[axis];
293 let mut mins = vec![f32::INFINITY; num_channels];
294 let mut maxs = vec![f32::NEG_INFINITY; num_channels];
295
296 for (i, &v) in data.iter().enumerate() {
297 let ch = channel_index(i, &shape, axis);
298 let f = v.to_f32().unwrap();
299 if f < mins[ch] {
300 mins[ch] = f;
301 }
302 if f > maxs[ch] {
303 maxs[ch] = f;
304 }
305 }
306
307 let params: Vec<(f32, i32)> = mins
308 .iter()
309 .zip(maxs.iter())
310 .map(|(&mn, &mx)| compute_scale_zp(mn, mx, dtype))
311 .collect();
312
313 let scales: Vec<f32> = params.iter().map(|&(s, _)| s).collect();
314 let zps: Vec<i32> = params.iter().map(|&(_, z)| z).collect();
315
316 let mut qdata = Vec::with_capacity(numel);
317 for (i, &v) in data.iter().enumerate() {
318 let ch = channel_index(i, &shape, axis);
319 qdata.push(quantize_val(
320 v.to_f32().unwrap(),
321 scales[ch],
322 zps[ch],
323 qmin,
324 qmax,
325 is_unsigned,
326 ));
327 }
328
329 Ok(QuantizedTensor {
330 data: qdata,
331 scale: scales,
332 zero_point: zps,
333 shape,
334 scheme,
335 dtype,
336 })
337 }
338 }
339}
340
341pub fn dequantize<T: Float>(qtensor: &QuantizedTensor) -> FerrotorchResult<Tensor<T>> {
349 let numel = qtensor.numel();
350 let mut result = Vec::with_capacity(numel);
351 let is_unsigned = qtensor.dtype == QuantDtype::Uint8;
352
353 match qtensor.scheme {
354 QuantScheme::PerTensor => {
355 let scale = qtensor.scale[0];
356 let zp = qtensor.zero_point[0];
357 for &q in &qtensor.data {
358 let val = (stored_to_i32(q, is_unsigned) - zp) as f32 * scale;
359 result.push(T::from(val).unwrap());
360 }
361 }
362 QuantScheme::PerChannel(axis) => {
363 for (i, &q) in qtensor.data.iter().enumerate() {
364 let ch = channel_index(i, &qtensor.shape, axis);
365 let val = (stored_to_i32(q, is_unsigned) - qtensor.zero_point[ch]) as f32
366 * qtensor.scale[ch];
367 result.push(T::from(val).unwrap());
368 }
369 }
370 }
371
372 Tensor::from_storage(TensorStorage::cpu(result), qtensor.shape.clone(), false)
373}
374
375pub fn quantized_matmul(
388 a: &QuantizedTensor,
389 b: &QuantizedTensor,
390) -> FerrotorchResult<QuantizedTensor> {
391 if a.shape.len() != 2 || b.shape.len() != 2 {
393 return Err(FerrotorchError::InvalidArgument {
394 message: format!(
395 "quantized_matmul requires 2-D tensors, got shapes {:?} and {:?}",
396 a.shape, b.shape
397 ),
398 });
399 }
400
401 let m = a.shape[0];
402 let k = a.shape[1];
403 let k2 = b.shape[0];
404 let n = b.shape[1];
405
406 if k != k2 {
407 return Err(FerrotorchError::ShapeMismatch {
408 message: format!(
409 "quantized_matmul inner dimensions mismatch: [{m}, {k}] x [{k2}, {n}]"
410 ),
411 });
412 }
413
414 if a.scale.len() != 1 || b.scale.len() != 1 {
416 return Err(FerrotorchError::InvalidArgument {
417 message: "quantized_matmul currently requires PerTensor-quantized inputs".into(),
418 });
419 }
420
421 let a_scale = a.scale[0];
422 let a_zp = a.zero_point[0];
423 let b_scale = b.scale[0];
424 let b_zp = b.zero_point[0];
425
426 let a_unsigned = a.dtype == QuantDtype::Uint8;
427 let b_unsigned = b.dtype == QuantDtype::Uint8;
428
429 let mut acc = vec![0i32; m * n];
431 for i in 0..m {
432 for j in 0..n {
433 let mut sum = 0i32;
434 for p in 0..k {
435 let qa = stored_to_i32(a.data[i * k + p], a_unsigned) - a_zp;
436 let qb = stored_to_i32(b.data[p * n + j], b_unsigned) - b_zp;
437 sum += qa * qb;
438 }
439 acc[i * n + j] = sum;
440 }
441 }
442
443 let combined_scale = a_scale * b_scale;
446
447 let mut out_min = f32::INFINITY;
449 let mut out_max = f32::NEG_INFINITY;
450 for &a_val in &acc {
451 let real = a_val as f32 * combined_scale;
452 if real < out_min {
453 out_min = real;
454 }
455 if real > out_max {
456 out_max = real;
457 }
458 }
459
460 let out_dtype = QuantDtype::Int8;
461 let (out_scale, out_zp) = compute_scale_zp(out_min, out_max, out_dtype);
462 let qmin = out_dtype.qmin();
463 let qmax = out_dtype.qmax();
464
465 let qdata: Vec<i8> = acc
466 .iter()
467 .map(|&a_val| {
468 let real = a_val as f32 * combined_scale;
469 quantize_val(real, out_scale, out_zp, qmin, qmax, false)
470 })
471 .collect();
472
473 Ok(QuantizedTensor {
474 data: qdata,
475 scale: vec![out_scale],
476 zero_point: vec![out_zp],
477 shape: vec![m, n],
478 scheme: QuantScheme::PerTensor,
479 dtype: out_dtype,
480 })
481}
482
483pub fn quantize_named_tensors<T: Float>(
494 named_tensors: impl IntoIterator<Item = (String, Tensor<T>)>,
495 scheme: QuantScheme,
496 dtype: QuantDtype,
497) -> FerrotorchResult<HashMap<String, QuantizedTensor>> {
498 let mut result = HashMap::new();
499 for (name, tensor) in named_tensors {
500 let qtensor = quantize(&tensor, scheme, dtype)?;
501 result.insert(name, qtensor);
502 }
503 Ok(result)
504}
505
506#[derive(Debug, Clone)]
512pub struct QParams {
513 pub scale: Vec<f32>,
515 pub zero_point: Vec<i32>,
517}
518
519impl QParams {
520 pub fn symmetric(max_abs: f32, dtype: QuantDtype) -> Self {
527 let max_abs = max_abs.max(f32::EPSILON);
528 match dtype {
529 QuantDtype::Int8 => QParams {
530 scale: vec![max_abs / 127.0],
531 zero_point: vec![0],
532 },
533 QuantDtype::Int4 => QParams {
534 scale: vec![max_abs / 7.0],
535 zero_point: vec![0],
536 },
537 QuantDtype::Uint8 => QParams {
538 scale: vec![max_abs / 128.0],
539 zero_point: vec![128],
540 },
541 }
542 }
543
544 pub fn asymmetric(min_val: f32, max_val: f32, dtype: QuantDtype) -> Self {
546 let (scale, zp) = compute_scale_zp(min_val, max_val, dtype);
547 QParams {
548 scale: vec![scale],
549 zero_point: vec![zp],
550 }
551 }
552}
553
554pub trait Observer {
560 fn observe(&mut self, data: &[f32]);
562 fn calculate_qparams(&self, dtype: QuantDtype) -> QParams;
564 fn reset(&mut self);
566}
567
568#[derive(Debug, Clone)]
576pub struct MinMaxObserver {
577 min_val: f32,
578 max_val: f32,
579}
580
581impl MinMaxObserver {
582 pub fn new() -> Self {
583 Self {
584 min_val: f32::INFINITY,
585 max_val: f32::NEG_INFINITY,
586 }
587 }
588}
589
590impl Default for MinMaxObserver {
591 fn default() -> Self {
592 Self::new()
593 }
594}
595
596impl Observer for MinMaxObserver {
597 fn observe(&mut self, data: &[f32]) {
598 for &x in data {
599 if !x.is_finite() {
600 continue;
601 }
602 if x < self.min_val {
603 self.min_val = x;
604 }
605 if x > self.max_val {
606 self.max_val = x;
607 }
608 }
609 }
610
611 fn calculate_qparams(&self, dtype: QuantDtype) -> QParams {
612 QParams::asymmetric(self.min_val, self.max_val, dtype)
613 }
614
615 fn reset(&mut self) {
616 self.min_val = f32::INFINITY;
617 self.max_val = f32::NEG_INFINITY;
618 }
619}
620
621#[derive(Debug, Clone)]
631pub struct PerChannelMinMaxObserver {
632 num_channels: usize,
633 axis: usize,
634 min_vals: Vec<f32>,
635 max_vals: Vec<f32>,
636}
637
638impl PerChannelMinMaxObserver {
639 pub fn new(num_channels: usize, axis: usize) -> Self {
644 Self {
645 num_channels,
646 axis,
647 min_vals: vec![f32::INFINITY; num_channels],
648 max_vals: vec![f32::NEG_INFINITY; num_channels],
649 }
650 }
651
652 pub fn observe_with_shape(&mut self, data: &[f32], shape: &[usize]) -> FerrotorchResult<()> {
656 if self.axis >= shape.len() {
657 return Err(FerrotorchError::InvalidArgument {
658 message: format!(
659 "PerChannelMinMaxObserver axis {} out of range for {}-d tensor",
660 self.axis,
661 shape.len()
662 ),
663 });
664 }
665 let actual_channels = shape[self.axis];
666 if actual_channels != self.num_channels {
667 return Err(FerrotorchError::InvalidArgument {
671 message: format!(
672 "PerChannelMinMaxObserver expected {} channels on axis {}, got {}",
673 self.num_channels, self.axis, actual_channels
674 ),
675 });
676 }
677
678 for (i, &x) in data.iter().enumerate() {
679 if !x.is_finite() {
680 continue;
681 }
682 let ch = channel_index(i, shape, self.axis);
683 if x < self.min_vals[ch] {
684 self.min_vals[ch] = x;
685 }
686 if x > self.max_vals[ch] {
687 self.max_vals[ch] = x;
688 }
689 }
690 Ok(())
691 }
692}
693
694impl Observer for PerChannelMinMaxObserver {
695 fn observe(&mut self, data: &[f32]) {
696 if !data.len().is_multiple_of(self.num_channels) {
702 return;
703 }
704 let per_channel = data.len() / self.num_channels;
705 for (i, &x) in data.iter().enumerate() {
706 if !x.is_finite() {
707 continue;
708 }
709 let ch = i / per_channel;
710 if ch >= self.num_channels {
711 continue;
712 }
713 if x < self.min_vals[ch] {
714 self.min_vals[ch] = x;
715 }
716 if x > self.max_vals[ch] {
717 self.max_vals[ch] = x;
718 }
719 }
720 }
721
722 fn calculate_qparams(&self, dtype: QuantDtype) -> QParams {
723 let params: Vec<(f32, i32)> = self
724 .min_vals
725 .iter()
726 .zip(self.max_vals.iter())
727 .map(|(&mn, &mx)| compute_scale_zp(mn, mx, dtype))
728 .collect();
729 QParams {
730 scale: params.iter().map(|&(s, _)| s).collect(),
731 zero_point: params.iter().map(|&(_, z)| z).collect(),
732 }
733 }
734
735 fn reset(&mut self) {
736 self.min_vals.fill(f32::INFINITY);
737 self.max_vals.fill(f32::NEG_INFINITY);
738 }
739}
740
741#[derive(Debug, Clone)]
750pub struct HistogramObserver {
751 num_bins: usize,
752 bins: Vec<u64>,
753 min_val: f32,
754 max_val: f32,
755 initialized: bool,
757}
758
759impl HistogramObserver {
760 pub fn new(num_bins: usize) -> Self {
761 Self {
762 num_bins,
763 bins: vec![0u64; num_bins],
764 min_val: f32::INFINITY,
765 max_val: f32::NEG_INFINITY,
766 initialized: false,
767 }
768 }
769
770 fn redistribute(&mut self, new_min: f32, new_max: f32) {
772 if !self.initialized || self.bins.iter().all(|&c| c == 0) {
773 self.min_val = new_min;
774 self.max_val = new_max;
775 return;
776 }
777
778 let old_min = self.min_val;
779 let old_max = self.max_val;
780 let old_range = old_max - old_min;
781 let new_range = new_max - new_min;
782
783 if old_range <= 0.0 || new_range <= 0.0 {
784 self.min_val = new_min;
785 self.max_val = new_max;
786 return;
787 }
788
789 let n = self.num_bins;
790 let old_bins = self.bins.clone();
791 self.bins.fill(0);
792
793 let old_bin_width = old_range / n as f32;
794 let new_bin_width = new_range / n as f32;
795
796 for (old_idx, &old_count) in old_bins.iter().enumerate().take(n) {
797 if old_count == 0 {
798 continue;
799 }
800 let old_center = old_min + (old_idx as f32 + 0.5) * old_bin_width;
802 let new_frac = (old_center - new_min) / new_bin_width;
804 let new_idx = (new_frac as usize).min(n - 1);
805 self.bins[new_idx] += old_count;
806 }
807
808 self.min_val = new_min;
809 self.max_val = new_max;
810 }
811}
812
813impl Observer for HistogramObserver {
814 fn observe(&mut self, data: &[f32]) {
815 let mut batch_min = f32::INFINITY;
817 let mut batch_max = f32::NEG_INFINITY;
818 for &x in data {
819 if !x.is_finite() {
820 continue;
821 }
822 if x < batch_min {
823 batch_min = x;
824 }
825 if x > batch_max {
826 batch_max = x;
827 }
828 }
829
830 if batch_min > batch_max {
831 return;
833 }
834
835 let new_min = if self.initialized {
837 self.min_val.min(batch_min)
838 } else {
839 batch_min
840 };
841 let new_max = if self.initialized {
842 self.max_val.max(batch_max)
843 } else {
844 batch_max
845 };
846
847 if self.initialized && (new_min < self.min_val || new_max > self.max_val) {
848 self.redistribute(new_min, new_max);
850 } else if !self.initialized {
851 self.min_val = new_min;
852 self.max_val = new_max;
853 self.initialized = true;
854 }
855
856 let range = (self.max_val - self.min_val).max(f32::EPSILON);
858 let n = self.num_bins;
859 for &x in data {
860 if !x.is_finite() {
861 continue;
862 }
863 let frac = (x - self.min_val) / range;
864 let idx = ((frac * n as f32) as usize).min(n - 1);
865 self.bins[idx] += 1;
866 }
867 }
868
869 fn calculate_qparams(&self, dtype: QuantDtype) -> QParams {
870 QParams::asymmetric(self.min_val, self.max_val, dtype)
871 }
872
873 fn reset(&mut self) {
874 self.bins.fill(0);
875 self.min_val = f32::INFINITY;
876 self.max_val = f32::NEG_INFINITY;
877 self.initialized = false;
878 }
879}
880
881#[derive(Debug, Clone)]
893pub struct FakeQuantize {
894 pub dtype: QuantDtype,
896 pub qparams: Option<QParams>,
898 pub observer_enabled: bool,
900 pub fake_quant_enabled: bool,
902 observer: MinMaxObserver,
904}
905
906impl FakeQuantize {
907 pub fn new(dtype: QuantDtype) -> Self {
909 Self {
910 dtype,
911 qparams: None,
912 observer_enabled: true,
913 fake_quant_enabled: true,
914 observer: MinMaxObserver::new(),
915 }
916 }
917
918 pub fn forward(&mut self, data: &[f32]) -> (Vec<f32>, Vec<f32>) {
923 if !self.fake_quant_enabled {
924 let ones = vec![1.0f32; data.len()];
925 return (data.to_vec(), ones);
926 }
927
928 if self.observer_enabled {
930 self.observer.observe(data);
931 }
932
933 let qparams = if let Some(cached) = self.qparams.as_ref().filter(|_| !self.observer_enabled)
936 {
937 cached.clone()
938 } else {
939 let qp = self.observer.calculate_qparams(self.dtype);
940 self.qparams = Some(qp.clone());
941 qp
942 };
943
944 let scale = qparams.scale[0];
945 let zp = qparams.zero_point[0];
946 let qmin = self.dtype.qmin();
947 let qmax = self.dtype.qmax();
948
949 let range_min = (qmin as f32 - zp as f32) * scale;
951 let range_max = (qmax as f32 - zp as f32) * scale;
952
953 let mut output = Vec::with_capacity(data.len());
954 let mut grad_mask = Vec::with_capacity(data.len());
955
956 for &x in data {
957 let q = (x / scale + zp as f32)
959 .round()
960 .clamp(qmin as f32, qmax as f32);
961 let dq = (q - zp as f32) * scale;
962 output.push(dq);
963
964 if x >= range_min && x <= range_max {
966 grad_mask.push(1.0);
967 } else {
968 grad_mask.push(0.0);
969 }
970 }
971
972 (output, grad_mask)
973 }
974}
975
976#[derive(Debug, Clone)]
982pub struct QatLayer {
983 pub weight_fq: FakeQuantize,
985 pub activation_fq: FakeQuantize,
987}
988
989#[derive(Debug)]
996pub struct QatModel {
997 pub layers: HashMap<String, QatLayer>,
999 pub dtype: QuantDtype,
1001}
1002
1003impl QatModel {
1004 pub fn new(dtype: QuantDtype) -> Self {
1006 Self {
1007 layers: HashMap::new(),
1008 dtype,
1009 }
1010 }
1011
1012 pub fn register_layer(&mut self, name: &str) {
1014 self.layers.insert(
1015 name.to_string(),
1016 QatLayer {
1017 weight_fq: FakeQuantize::new(self.dtype),
1018 activation_fq: FakeQuantize::new(self.dtype),
1019 },
1020 );
1021 }
1022
1023 pub fn fake_quantize_weights(
1028 &mut self,
1029 layer_name: &str,
1030 weights: &[f32],
1031 ) -> FerrotorchResult<(Vec<f32>, Vec<f32>)> {
1032 let layer =
1033 self.layers
1034 .get_mut(layer_name)
1035 .ok_or_else(|| FerrotorchError::InvalidArgument {
1036 message: format!("layer '{layer_name}' not registered for QAT"),
1037 })?;
1038
1039 let originals = weights.to_vec();
1041
1042 let (fq_weights, _mask) = layer.weight_fq.forward(weights);
1044
1045 Ok((fq_weights, originals))
1046 }
1047
1048 pub fn fake_quantize_activations(
1052 &mut self,
1053 layer_name: &str,
1054 activations: &[f32],
1055 ) -> FerrotorchResult<(Vec<f32>, Vec<f32>)> {
1056 let layer =
1057 self.layers
1058 .get_mut(layer_name)
1059 .ok_or_else(|| FerrotorchError::InvalidArgument {
1060 message: format!("layer '{layer_name}' not registered for QAT"),
1061 })?;
1062
1063 let (fq_activations, grad_mask) = layer.activation_fq.forward(activations);
1064 Ok((fq_activations, grad_mask))
1065 }
1066}
1067
1068pub fn prepare_qat(param_names: &[&str], dtype: QuantDtype) -> QatModel {
1073 let mut model = QatModel::new(dtype);
1074
1075 for &name in param_names {
1076 let layer_name = if let Some(prefix) = name.strip_suffix(".weight") {
1078 prefix
1079 } else if let Some(prefix) = name.strip_suffix(".bias") {
1080 if !model.layers.contains_key(prefix) {
1083 model.register_layer(prefix);
1084 }
1085 continue;
1086 } else {
1087 name
1088 };
1089
1090 model.register_layer(layer_name);
1091 }
1092
1093 model
1094}
1095
1096pub mod cuda_rng {
1105 use std::sync::Mutex;
1106
1107 static RNG_STATE: Mutex<u64> = Mutex::new(0xdeadbeef_cafebabe);
1109
1110 static RNG_STACK: Mutex<Vec<u64>> = Mutex::new(Vec::new());
1112
1113 pub fn get_state() -> u64 {
1115 let guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1116 *guard
1117 }
1118
1119 pub fn set_state(state: u64) {
1121 let mut guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1122 *guard = state;
1123 }
1124
1125 pub fn fork_rng(new_seed: u64) {
1130 let current = {
1131 let guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1132 *guard
1133 };
1134
1135 {
1136 let mut stack = RNG_STACK.lock().unwrap_or_else(|e| e.into_inner());
1137 stack.push(current);
1138 }
1139
1140 set_state(new_seed);
1141 }
1142
1143 pub fn join_rng() {
1148 let saved = {
1149 let mut stack = RNG_STACK.lock().unwrap_or_else(|e| e.into_inner());
1150 stack.pop()
1151 };
1152
1153 if let Some(state) = saved {
1154 set_state(state);
1155 }
1156 }
1157
1158 pub fn next_seed() -> u64 {
1160 let mut guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1161 *guard = guard.wrapping_add(0x9e3779b97f4a7c15);
1163 let mut z = *guard;
1164 z = (z ^ (z >> 30)).wrapping_mul(0xbf58476d1ce4e5b9);
1165 z = (z ^ (z >> 27)).wrapping_mul(0x94d049bb133111eb);
1166 z ^ (z >> 31)
1167 }
1168}
1169
1170#[cfg(test)]
1175mod tests {
1176 use super::*;
1177
1178 fn make_tensor(data: &[f32], shape: &[usize]) -> Tensor<f32> {
1180 crate::from_slice(data, shape).unwrap()
1181 }
1182
1183 #[test]
1186 fn test_per_tensor_int8_roundtrip() {
1187 let data: Vec<f32> = (-10..=10).map(|x| x as f32 * 0.5).collect();
1188 let t = make_tensor(&data, &[data.len()]);
1189 let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1190 let rt: Tensor<f32> = dequantize(&qt).unwrap();
1191
1192 assert_eq!(rt.shape(), t.shape());
1193 let orig = t.data().unwrap();
1194 let recovered = rt.data().unwrap();
1195 for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1196 let err = (o - r).abs();
1197 assert!(
1199 err < 0.05,
1200 "element {i}: original={o}, recovered={r}, error={err}"
1201 );
1202 }
1203 }
1204
1205 #[test]
1206 fn test_per_tensor_uint8_roundtrip() {
1207 let data: Vec<f32> = (0..=20).map(|x| x as f32 * 0.1).collect();
1208 let t = make_tensor(&data, &[data.len()]);
1209 let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Uint8).unwrap();
1210 let rt: Tensor<f32> = dequantize(&qt).unwrap();
1211
1212 let orig = t.data().unwrap();
1213 let recovered = rt.data().unwrap();
1214 for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1215 let err = (o - r).abs();
1216 assert!(
1218 err < 0.02,
1219 "element {i}: original={o}, recovered={r}, error={err}"
1220 );
1221 }
1222 }
1223
1224 #[test]
1225 fn test_per_tensor_int4_roundtrip() {
1226 let data: Vec<f32> = (-8..=7).map(|x| x as f32).collect();
1228 let t = make_tensor(&data, &[data.len()]);
1229 let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int4).unwrap();
1230 let rt: Tensor<f32> = dequantize(&qt).unwrap();
1231
1232 let orig = t.data().unwrap();
1233 let recovered = rt.data().unwrap();
1234 for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1235 let err = (o - r).abs();
1236 assert!(
1238 err < 1.01,
1239 "element {i}: original={o}, recovered={r}, error={err}"
1240 );
1241 }
1242 }
1243
1244 #[test]
1247 fn test_per_channel_int8_roundtrip() {
1248 #[rustfmt::skip]
1250 let data: Vec<f32> = vec![
1251 0.0, 1.0, 2.0, 3.0,
1253 -10.0, -5.0, 5.0, 10.0,
1255 100.0, 130.0, 170.0, 200.0,
1257 ];
1258 let t = make_tensor(&data, &[3, 4]);
1259 let qt = quantize(&t, QuantScheme::PerChannel(0), QuantDtype::Int8).unwrap();
1260 let rt: Tensor<f32> = dequantize(&qt).unwrap();
1261
1262 assert_eq!(qt.scale.len(), 3);
1263 assert_eq!(qt.zero_point.len(), 3);
1264
1265 let orig = t.data().unwrap();
1266 let recovered = rt.data().unwrap();
1267 for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1268 let err = (o - r).abs();
1269 assert!(
1272 err < 0.5,
1273 "element {i}: original={o}, recovered={r}, error={err}"
1274 );
1275 }
1276 }
1277
1278 #[test]
1279 fn test_per_channel_axis_out_of_bounds() {
1280 let t = make_tensor(&[1.0, 2.0, 3.0], &[3]);
1281 let result = quantize(&t, QuantScheme::PerChannel(5), QuantDtype::Int8);
1282 assert!(result.is_err());
1283 }
1284
1285 #[test]
1288 fn test_quantized_matmul_identity() {
1289 let a_data = vec![1.0f32, 2.0, 3.0, 4.0];
1291 let a = make_tensor(&a_data, &[2, 2]);
1292 let eye = make_tensor(&[1.0, 0.0, 0.0, 1.0], &[2, 2]);
1293
1294 let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1295 let qi = quantize(&eye, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1296 let qc = quantized_matmul(&qa, &qi).unwrap();
1297 let c: Tensor<f32> = dequantize(&qc).unwrap();
1298
1299 assert_eq!(c.shape(), &[2, 2]);
1300 let c_data = c.data().unwrap();
1301 for (i, (&expected, &got)) in a_data.iter().zip(c_data.iter()).enumerate() {
1302 let err = (expected - got).abs();
1303 assert!(
1304 err < 0.5,
1305 "element {i}: expected={expected}, got={got}, error={err}"
1306 );
1307 }
1308 }
1309
1310 #[test]
1311 fn test_quantized_matmul_correctness() {
1312 let a = make_tensor(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
1321 let b = make_tensor(&[7.0, 8.0, 9.0, 10.0, 11.0, 12.0], &[3, 2]);
1322
1323 let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1324 let qb = quantize(&b, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1325 let qc = quantized_matmul(&qa, &qb).unwrap();
1326 let c: Tensor<f32> = dequantize(&qc).unwrap();
1327
1328 let expected = [58.0f32, 64.0, 139.0, 154.0];
1329 let c_data = c.data().unwrap();
1330 assert_eq!(c.shape(), &[2, 2]);
1331 for (i, (&e, &g)) in expected.iter().zip(c_data.iter()).enumerate() {
1332 let err = (e - g).abs();
1333 assert!(err < 3.0, "element {i}: expected={e}, got={g}, error={err}");
1336 }
1337 }
1338
1339 #[test]
1340 fn test_quantized_matmul_shape_mismatch() {
1341 let a = make_tensor(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
1342 let b = make_tensor(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
1343
1344 let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1345 let qb = quantize(&b, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1346 let result = quantized_matmul(&qa, &qb);
1347 assert!(result.is_err());
1348 }
1349
1350 #[test]
1351 fn test_quantized_matmul_non_2d() {
1352 let a = make_tensor(&[1.0, 2.0, 3.0], &[3]);
1353 let b = make_tensor(&[4.0, 5.0, 6.0], &[3]);
1354
1355 let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1356 let qb = quantize(&b, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1357 let result = quantized_matmul(&qa, &qb);
1358 assert!(result.is_err());
1359 }
1360
1361 #[test]
1364 fn test_quantize_named_tensors() {
1365 let w1 = make_tensor(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
1366 let w2 = make_tensor(&[-1.0, 0.0, 1.0, 2.0, 3.0, 4.0], &[3, 2]);
1367
1368 let named = vec![
1369 ("layer.weight".to_string(), w1),
1370 ("layer2.weight".to_string(), w2),
1371 ];
1372
1373 let qmap = quantize_named_tensors(named, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1374
1375 assert_eq!(qmap.len(), 2);
1376 assert!(qmap.contains_key("layer.weight"));
1377 assert!(qmap.contains_key("layer2.weight"));
1378 assert_eq!(qmap["layer.weight"].shape(), &[2, 2]);
1379 assert_eq!(qmap["layer2.weight"].shape(), &[3, 2]);
1380 }
1381
1382 #[test]
1385 fn test_quantize_constant_tensor() {
1386 let t = make_tensor(&[5.0, 5.0, 5.0, 5.0], &[4]);
1388 let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1389 let rt: Tensor<f32> = dequantize(&qt).unwrap();
1390
1391 let recovered = rt.data().unwrap();
1392 for &r in recovered {
1393 assert!(
1394 (r - 5.0).abs() < 0.1,
1395 "constant tensor dequantized to {r}, expected 5.0"
1396 );
1397 }
1398 }
1399
1400 #[test]
1401 fn test_quantize_single_element() {
1402 let t = make_tensor(&[42.0], &[1]);
1403 let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1404 let rt: Tensor<f32> = dequantize(&qt).unwrap();
1405 assert!((rt.data().unwrap()[0] - 42.0).abs() < 0.5);
1406 }
1407
1408 #[test]
1409 fn test_per_channel_int4() {
1410 let data = vec![0.0, 1.0, 2.0, -4.0, 0.0, 4.0];
1412 let t = make_tensor(&data, &[2, 3]);
1413 let qt = quantize(&t, QuantScheme::PerChannel(0), QuantDtype::Int4).unwrap();
1414
1415 assert_eq!(qt.scale.len(), 2);
1416 assert_eq!(qt.zero_point.len(), 2);
1417
1418 let rt: Tensor<f32> = dequantize(&qt).unwrap();
1419 let orig = t.data().unwrap();
1420 let recovered = rt.data().unwrap();
1421 for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1422 let err = (o - r).abs();
1423 assert!(
1425 err < 1.0,
1426 "element {i}: original={o}, recovered={r}, error={err}"
1427 );
1428 }
1429 }
1430
1431 #[test]
1432 fn test_dequantize_f64() {
1433 let data = vec![1.0f32, 2.0, 3.0, 4.0];
1434 let t = crate::from_slice(&data, &[4]).unwrap();
1435 let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1436 let rt: Tensor<f64> = dequantize(&qt).unwrap();
1437
1438 assert_eq!(rt.shape(), &[4]);
1439 let recovered = rt.data().unwrap();
1440 for (i, &r) in recovered.iter().enumerate() {
1441 let expected = data[i] as f64;
1442 let err = (expected - r).abs();
1443 assert!(
1444 err < 0.05,
1445 "element {i}: expected={expected}, recovered={r}, error={err}"
1446 );
1447 }
1448 }
1449
1450 #[test]
1451 fn test_quantized_tensor_accessors() {
1452 let t = make_tensor(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
1453 let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1454
1455 assert_eq!(qt.numel(), 6);
1456 assert_eq!(qt.shape(), &[2, 3]);
1457 assert_eq!(qt.data().len(), 6);
1458 assert_eq!(qt.scale().len(), 1);
1459 assert_eq!(qt.zero_point().len(), 1);
1460 assert_eq!(qt.scheme(), QuantScheme::PerTensor);
1461 assert_eq!(qt.qdtype(), QuantDtype::Int8);
1462 }
1463
1464 #[test]
1467 fn test_qparams_symmetric_int8() {
1468 let qp = QParams::symmetric(5.0, QuantDtype::Int8);
1469 assert_eq!(qp.zero_point, vec![0]);
1470 assert!((qp.scale[0] - 5.0 / 127.0).abs() < 1e-7);
1471 }
1472
1473 #[test]
1474 fn test_qparams_symmetric_uint8() {
1475 let qp = QParams::symmetric(5.0, QuantDtype::Uint8);
1476 assert_eq!(qp.zero_point, vec![128]);
1477 assert!((qp.scale[0] - 5.0 / 128.0).abs() < 1e-7);
1478 }
1479
1480 #[test]
1481 fn test_qparams_symmetric_int4() {
1482 let qp = QParams::symmetric(7.0, QuantDtype::Int4);
1483 assert_eq!(qp.zero_point, vec![0]);
1484 assert!((qp.scale[0] - 1.0).abs() < 1e-7);
1485 }
1486
1487 #[test]
1490 fn test_minmax_observer() {
1491 let mut obs = MinMaxObserver::new();
1492 obs.observe(&[1.0, 2.0, 3.0]);
1493 obs.observe(&[-1.0, 5.0]);
1494 let qp = obs.calculate_qparams(QuantDtype::Int8);
1495 assert_eq!(qp.scale.len(), 1);
1497 assert_eq!(qp.zero_point.len(), 1);
1498 }
1499
1500 #[test]
1501 fn test_minmax_observer_filters_nan_inf() {
1502 let mut obs = MinMaxObserver::new();
1503 obs.observe(&[1.0, f32::NAN, 2.0, f32::INFINITY, -1.0, f32::NEG_INFINITY]);
1504 let qp = obs.calculate_qparams(QuantDtype::Int8);
1505 let expected_range = 2.0 - (-1.0); let expected_scale = expected_range / 255.0;
1508 assert!((qp.scale[0] - expected_scale).abs() < 1e-5);
1509 }
1510
1511 #[test]
1514 fn test_per_channel_observer_with_shape() {
1515 let mut obs = PerChannelMinMaxObserver::new(2, 0);
1516 obs.observe_with_shape(&[0.0, 1.0, 2.0, 10.0, 20.0, 30.0], &[2, 3])
1518 .unwrap();
1519 let qp = obs.calculate_qparams(QuantDtype::Int8);
1520 assert_eq!(qp.scale.len(), 2);
1521 assert_eq!(qp.zero_point.len(), 2);
1522 }
1523
1524 #[test]
1525 fn test_per_channel_observer_shape_mismatch() {
1526 let mut obs = PerChannelMinMaxObserver::new(3, 0);
1527 let result = obs.observe_with_shape(&[1.0; 6], &[2, 3]);
1529 assert!(result.is_err());
1530 }
1531
1532 #[test]
1533 fn test_per_channel_observer_axis() {
1534 let mut obs = PerChannelMinMaxObserver::new(3, 1);
1535 obs.observe_with_shape(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3])
1537 .unwrap();
1538 let qp = obs.calculate_qparams(QuantDtype::Int8);
1539 assert_eq!(qp.scale.len(), 3);
1540 }
1541
1542 #[test]
1543 fn test_per_channel_observer_filters_nan_inf() {
1544 let mut obs = PerChannelMinMaxObserver::new(2, 0);
1545 obs.observe_with_shape(&[f32::NAN, 1.0, 2.0, 10.0, f32::INFINITY, 30.0], &[2, 3])
1546 .unwrap();
1547 let qp = obs.calculate_qparams(QuantDtype::Int8);
1549 assert_eq!(qp.scale.len(), 2);
1550 }
1551
1552 #[test]
1555 fn test_histogram_observer_basic() {
1556 let mut obs = HistogramObserver::new(100);
1557 obs.observe(&[0.0, 0.5, 1.0]);
1558 let qp = obs.calculate_qparams(QuantDtype::Int8);
1559 assert_eq!(qp.scale.len(), 1);
1560 }
1561
1562 #[test]
1563 fn test_histogram_observer_range_expansion() {
1564 let mut obs = HistogramObserver::new(100);
1565 obs.observe(&[0.0, 1.0]);
1566 let bins_after_first = obs.bins.clone();
1568 let total_first: u64 = bins_after_first.iter().sum();
1569 assert_eq!(total_first, 2);
1570
1571 obs.observe(&[-1.0, 2.0]);
1572 let total_second: u64 = obs.bins.iter().sum();
1574 assert_eq!(total_second, 4);
1576 }
1577
1578 #[test]
1579 fn test_histogram_observer_filters_nan_inf() {
1580 let mut obs = HistogramObserver::new(50);
1581 obs.observe(&[f32::NAN, 1.0, f32::INFINITY, 2.0]);
1582 let total: u64 = obs.bins.iter().sum();
1583 assert_eq!(total, 2);
1585 }
1586
1587 #[test]
1590 fn test_fake_quantize_roundtrip() {
1591 let mut fq = FakeQuantize::new(QuantDtype::Int8);
1592 let data = vec![0.0, 0.5, 1.0, 1.5, 2.0];
1593 let (output, mask) = fq.forward(&data);
1594 assert_eq!(output.len(), 5);
1595 assert_eq!(mask.len(), 5);
1596
1597 for (i, (&o, &d)) in output.iter().zip(data.iter()).enumerate() {
1599 assert!((o - d).abs() < 0.1, "element {i}: output={o}, data={d}");
1600 }
1601 }
1602
1603 #[test]
1604 #[allow(clippy::float_cmp)]
1607 fn test_fake_quantize_ste_clipping() {
1608 let mut fq = FakeQuantize::new(QuantDtype::Int8);
1609 let (_, _) = fq.forward(&[0.0, 1.0, 2.0]);
1611
1612 fq.observer_enabled = false;
1614
1615 let (_, mask) = fq.forward(&[0.5, 1.0, 100.0, -100.0]);
1617 assert_eq!(mask[0], 1.0);
1619 assert_eq!(mask[1], 1.0);
1620 assert_eq!(mask[2], 0.0);
1622 assert_eq!(mask[3], 0.0);
1623 }
1624
1625 #[test]
1626 fn test_fake_quantize_observer_disabled_uses_cached() {
1627 let mut fq = FakeQuantize::new(QuantDtype::Int8);
1628 let (_, _) = fq.forward(&[0.0, 10.0]);
1630 let cached_scale = fq.qparams.as_ref().unwrap().scale[0];
1631
1632 fq.observer_enabled = false;
1634
1635 let (_, _) = fq.forward(&[0.0, 1000.0]);
1637 let scale_after = fq.qparams.as_ref().unwrap().scale[0];
1638 assert!(
1639 (scale_after - cached_scale).abs() < 1e-10,
1640 "scale should not change when observer is disabled"
1641 );
1642 }
1643
1644 #[test]
1645 #[allow(clippy::float_cmp)]
1648 fn test_fake_quantize_disabled_is_identity() {
1649 let mut fq = FakeQuantize::new(QuantDtype::Int8);
1650 fq.fake_quant_enabled = false;
1651 let data = vec![1.234, 5.678, -9.012];
1652 let (output, mask) = fq.forward(&data);
1653 assert_eq!(output, data);
1654 assert!(mask.iter().all(|&m| m == 1.0));
1655 }
1656
1657 #[test]
1660 fn test_qat_model_register_and_fq_weights() {
1661 let mut model = QatModel::new(QuantDtype::Int8);
1662 model.register_layer("fc1");
1663
1664 let weights = vec![0.1, 0.2, 0.3, 0.4];
1665 let (fq_weights, originals) = model.fake_quantize_weights("fc1", &weights).unwrap();
1666
1667 assert_eq!(originals, weights);
1669 for (i, (&fq, &orig)) in fq_weights.iter().zip(weights.iter()).enumerate() {
1671 assert!((fq - orig).abs() < 0.1, "weight {i}: fq={fq}, orig={orig}");
1672 }
1673 }
1674
1675 #[test]
1676 fn test_qat_model_activation_fq_per_layer() {
1677 let mut model = QatModel::new(QuantDtype::Int8);
1678 model.register_layer("layer1");
1679 model.register_layer("layer2");
1680
1681 let (act1, _) = model
1683 .fake_quantize_activations("layer1", &[1.0, 2.0])
1684 .unwrap();
1685 let (act2, _) = model
1686 .fake_quantize_activations("layer2", &[10.0, 20.0])
1687 .unwrap();
1688 assert_eq!(act1.len(), 2);
1689 assert_eq!(act2.len(), 2);
1690 }
1691
1692 #[test]
1693 fn test_qat_model_unregistered_layer_errors() {
1694 let mut model = QatModel::new(QuantDtype::Int8);
1695 let result = model.fake_quantize_weights("nonexistent", &[1.0]);
1696 assert!(result.is_err());
1697 }
1698
1699 #[test]
1702 fn test_prepare_qat_skips_bias() {
1703 let names = &["fc1.weight", "fc1.bias", "fc2.weight", "fc2.bias"];
1704 let model = prepare_qat(names, QuantDtype::Int8);
1705
1706 assert!(model.layers.contains_key("fc1"));
1707 assert!(model.layers.contains_key("fc2"));
1708 assert_eq!(model.layers.len(), 2);
1709 }
1710
1711 #[test]
1712 fn test_prepare_qat_bias_only_still_registers() {
1713 let names = &["fc1.bias"];
1714 let model = prepare_qat(names, QuantDtype::Int8);
1715 assert!(model.layers.contains_key("fc1"));
1717 }
1718
1719 fn cuda_rng_test_lock() -> std::sync::MutexGuard<'static, ()> {
1727 static LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
1728 LOCK.lock().unwrap_or_else(|p| p.into_inner())
1729 }
1730
1731 #[test]
1732 fn test_cuda_rng_fork_join() {
1733 let _g = cuda_rng_test_lock();
1734 let initial = cuda_rng::get_state();
1735 cuda_rng::fork_rng(0x12345678);
1736 assert_eq!(cuda_rng::get_state(), 0x12345678);
1737 cuda_rng::join_rng();
1738 assert_eq!(cuda_rng::get_state(), initial);
1739 }
1740
1741 #[test]
1742 fn test_cuda_rng_next_seed() {
1743 let _g = cuda_rng_test_lock();
1744 let s1 = cuda_rng::next_seed();
1745 let s2 = cuda_rng::next_seed();
1746 assert_ne!(s1, s2, "consecutive seeds should differ");
1747 }
1748}