1use crate::error::{ErrorBuilder, ErrorCode, FfiError};
114use serde::{Deserialize, Serialize};
115
116#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
118pub enum QuantizationType {
119 Int8,
121 Int4,
123 Uint8,
125 Fp16,
127 BFloat16,
129 Dynamic,
131}
132
133impl QuantizationType {
134 pub fn bits(&self) -> u8 {
136 match self {
137 QuantizationType::Int4 => 4,
138 QuantizationType::Int8 | QuantizationType::Uint8 => 8,
139 QuantizationType::Fp16 | QuantizationType::BFloat16 => 16,
140 QuantizationType::Dynamic => 8, }
142 }
143
144 pub fn compression_ratio(&self) -> f32 {
146 32.0 / self.bits() as f32
147 }
148
149 pub fn value_range(&self) -> (f64, f64) {
151 match self {
152 QuantizationType::Int8 => (-128.0, 127.0),
153 QuantizationType::Int4 => (-8.0, 7.0),
154 QuantizationType::Uint8 => (0.0, 255.0),
155 QuantizationType::Fp16 => (-65504.0, 65504.0),
156 QuantizationType::BFloat16 => (-3.39e38, 3.39e38),
157 QuantizationType::Dynamic => (-128.0, 127.0),
158 }
159 }
160}
161
162#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
164pub enum QuantizationGranularity {
165 PerTensor,
167 PerChannel,
169 PerGroup { group_size: usize },
171}
172
173#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
175pub enum QuantizationScheme {
176 Symmetric,
178 Asymmetric,
180}
181
182#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
184pub enum CalibrationMethod {
185 MinMax,
187 Percentile { lower: f32, upper: f32 },
189 MovingAverage { momentum: f32 },
191 Mse,
193 Entropy,
195}
196
197impl Default for CalibrationMethod {
198 fn default() -> Self {
199 CalibrationMethod::Percentile {
200 lower: 0.01,
201 upper: 99.99,
202 }
203 }
204}
205
206#[derive(Debug, Clone, Serialize, Deserialize)]
208pub struct QuantizationParams {
209 pub scale: f32,
211 pub zero_point: i32,
213 pub qtype: QuantizationType,
215 pub range: (f32, f32),
217}
218
219impl QuantizationParams {
220 pub fn new(scale: f32, zero_point: i32, qtype: QuantizationType) -> Self {
227 Self {
228 scale,
229 zero_point,
230 qtype,
231 range: (0.0, 0.0),
232 }
233 }
234
235 pub fn from_range(
243 min_val: f32,
244 max_val: f32,
245 qtype: QuantizationType,
246 scheme: QuantizationScheme,
247 ) -> Self {
248 let (qmin, qmax) = qtype.value_range();
249
250 let (scale, zero_point) = match scheme {
251 QuantizationScheme::Symmetric => {
252 let max_abs = min_val.abs().max(max_val.abs());
254 let scale = (2.0 * max_abs) / (qmax - qmin) as f32;
255 (scale, 0)
256 }
257 QuantizationScheme::Asymmetric => {
258 let scale = (max_val - min_val) / (qmax - qmin) as f32;
260 let zero_point = qmin as f32 - min_val / scale;
261 (scale, zero_point.round() as i32)
262 }
263 };
264
265 Self {
266 scale: scale.max(1e-8), zero_point,
268 qtype,
269 range: (min_val, max_val),
270 }
271 }
272
273 pub fn quantize(&self, value: f32) -> i32 {
275 let (qmin, qmax) = self.qtype.value_range();
276 let quantized = (value / self.scale).round() + self.zero_point as f32;
277 quantized.clamp(qmin as f32, qmax as f32) as i32
278 }
279
280 pub fn dequantize(&self, quantized: i32) -> f32 {
282 (quantized - self.zero_point) as f32 * self.scale
283 }
284
285 pub fn quantize_array(&self, values: &[f32]) -> Vec<i32> {
287 values.iter().map(|&v| self.quantize(v)).collect()
288 }
289
290 pub fn dequantize_array(&self, quantized: &[i32]) -> Vec<f32> {
292 quantized.iter().map(|&q| self.dequantize(q)).collect()
293 }
294}
295
296#[derive(Debug, Clone, Serialize, Deserialize)]
298pub struct QuantizationConfig {
299 pub qtype: QuantizationType,
301 pub granularity: QuantizationGranularity,
303 pub scheme: QuantizationScheme,
305 pub calibration: CalibrationMethod,
307 pub quantize_weights: bool,
309 pub quantize_activations: bool,
311 pub quantize_biases: bool,
313 pub skip_layers: Vec<String>,
315 pub force_fp32_layers: Vec<String>,
317}
318
319impl QuantizationConfig {
320 pub fn new(qtype: QuantizationType) -> Self {
322 Self {
323 qtype,
324 granularity: QuantizationGranularity::PerChannel,
325 scheme: QuantizationScheme::Asymmetric,
326 calibration: CalibrationMethod::default(),
327 quantize_weights: true,
328 quantize_activations: true,
329 quantize_biases: false, skip_layers: Vec::new(),
331 force_fp32_layers: vec!["output".to_string()], }
333 }
334
335 pub fn with_granularity(mut self, granularity: QuantizationGranularity) -> Self {
337 self.granularity = granularity;
338 self
339 }
340
341 pub fn with_scheme(mut self, scheme: QuantizationScheme) -> Self {
343 self.scheme = scheme;
344 self
345 }
346
347 pub fn with_calibration(mut self, calibration: CalibrationMethod) -> Self {
349 self.calibration = calibration;
350 self
351 }
352
353 pub fn skip_layer(mut self, layer_name: String) -> Self {
355 self.skip_layers.push(layer_name);
356 self
357 }
358
359 pub fn force_fp32_layer(mut self, layer_name: String) -> Self {
361 self.force_fp32_layers.push(layer_name);
362 self
363 }
364}
365
366impl Default for QuantizationConfig {
367 fn default() -> Self {
368 Self::new(QuantizationType::Int8)
369 }
370}
371
372#[derive(Debug, Clone, Serialize, Deserialize)]
374pub struct QuantizedTensor {
375 pub quantized_data: Vec<i32>,
377 pub params: QuantizationParams,
379 pub shape: Vec<usize>,
381 pub name: Option<String>,
383}
384
385impl QuantizedTensor {
386 pub fn new(quantized_data: Vec<i32>, params: QuantizationParams, shape: Vec<usize>) -> Self {
388 Self {
389 quantized_data,
390 params,
391 shape,
392 name: None,
393 }
394 }
395
396 pub fn with_name(mut self, name: String) -> Self {
398 self.name = Some(name);
399 self
400 }
401
402 pub fn size_bytes(&self) -> usize {
404 let bits_per_element = self.params.qtype.bits() as usize;
405 (self.quantized_data.len() * bits_per_element + 7) / 8
406 }
407
408 pub fn compression_ratio(&self) -> f32 {
410 let fp32_size = self.quantized_data.len() * 4; let quantized_size = self.size_bytes();
412 fp32_size as f32 / quantized_size as f32
413 }
414
415 pub fn dequantize(&self) -> Vec<f32> {
417 self.params.dequantize_array(&self.quantized_data)
418 }
419
420 pub fn quantization_error(&self, original: &[f32]) -> QuantizationError {
422 let dequantized = self.dequantize();
423
424 let mut mse = 0.0_f32;
425 let mut max_error = 0.0_f32;
426
427 for (orig, deq) in original.iter().zip(dequantized.iter()) {
428 let error = (orig - deq).abs();
429 mse += error * error;
430 max_error = max_error.max(error);
431 }
432
433 mse /= original.len() as f32;
434 let rmse = mse.sqrt();
435
436 let signal_power: f32 = original.iter().map(|x| x * x).sum::<f32>() / original.len() as f32;
438 let sqnr_db = 10.0 * (signal_power / mse).log10();
439
440 QuantizationError {
441 mse,
442 rmse,
443 max_error,
444 sqnr_db,
445 }
446 }
447}
448
449#[derive(Debug, Clone, Serialize, Deserialize)]
451pub struct QuantizationError {
452 pub mse: f32,
454 pub rmse: f32,
456 pub max_error: f32,
458 pub sqnr_db: f32,
460}
461
462#[derive(Debug, Clone)]
464pub struct CalibrationDataset {
465 samples: Vec<Vec<f32>>,
467 max_samples: usize,
469}
470
471impl CalibrationDataset {
472 pub fn new(max_samples: usize) -> Self {
474 Self {
475 samples: Vec::new(),
476 max_samples,
477 }
478 }
479
480 pub fn add_sample(&mut self, sample: Vec<f32>) {
482 if self.samples.len() < self.max_samples {
483 self.samples.push(sample);
484 }
485 }
486
487 pub fn samples(&self) -> &[Vec<f32>] {
489 &self.samples
490 }
491
492 pub fn compute_statistics(&self) -> CalibrationStatistics {
494 let mut all_values: Vec<f32> = self
495 .samples
496 .iter()
497 .flat_map(|s| s.iter().copied())
498 .collect();
499
500 all_values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
501
502 let min_val = all_values.first().copied().unwrap_or(0.0);
503 let max_val = all_values.last().copied().unwrap_or(0.0);
504
505 let mean = all_values.iter().sum::<f32>() / all_values.len() as f32;
506
507 let variance: f32 =
508 all_values.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / all_values.len() as f32;
509 let std_dev = variance.sqrt();
510
511 CalibrationStatistics {
512 min_val,
513 max_val,
514 mean,
515 std_dev,
516 num_samples: all_values.len(),
517 }
518 }
519
520 pub fn percentile(&self, percentile: f32) -> f32 {
522 let mut all_values: Vec<f32> = self
523 .samples
524 .iter()
525 .flat_map(|s| s.iter().copied())
526 .collect();
527 all_values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
528
529 let index = ((percentile / 100.0) * all_values.len() as f32) as usize;
530 all_values
531 .get(index.min(all_values.len() - 1))
532 .copied()
533 .unwrap_or(0.0)
534 }
535}
536
537#[derive(Debug, Clone, Serialize, Deserialize)]
539pub struct CalibrationStatistics {
540 pub min_val: f32,
541 pub max_val: f32,
542 pub mean: f32,
543 pub std_dev: f32,
544 pub num_samples: usize,
545}
546
547#[derive(Debug, Clone)]
549pub struct Quantizer {
550 config: QuantizationConfig,
551 calibration_data: Option<CalibrationDataset>,
552}
553
554impl Quantizer {
555 pub fn new(config: QuantizationConfig) -> Self {
557 Self {
558 config,
559 calibration_data: None,
560 }
561 }
562
563 pub fn with_calibration_data(mut self, data: CalibrationDataset) -> Self {
565 self.calibration_data = Some(data);
566 self
567 }
568
569 pub fn quantize_tensor(
576 &self,
577 data: &[f32],
578 shape: Vec<usize>,
579 name: Option<String>,
580 ) -> Result<QuantizedTensor, FfiError> {
581 if let Some(ref n) = name {
583 if self.config.skip_layers.contains(n) || self.config.force_fp32_layers.contains(n) {
584 return Err(FfiError::Enhanced(
585 ErrorBuilder::new(ErrorCode::OperationFailed)
586 .message("Layer marked for FP32 preservation")
587 .context("layer", n)
588 .build(),
589 ));
590 }
591 }
592
593 let (min_val, max_val) = match &self.calibration_data {
595 Some(calib) => match self.config.calibration {
596 CalibrationMethod::MinMax => {
597 let stats = calib.compute_statistics();
598 (stats.min_val, stats.max_val)
599 }
600 CalibrationMethod::Percentile { lower, upper } => {
601 (calib.percentile(lower), calib.percentile(upper))
602 }
603 _ => {
604 let min = data.iter().fold(f32::INFINITY, |a, &b| a.min(b));
606 let max = data.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
607 (min, max)
608 }
609 },
610 None => {
611 let min = data.iter().fold(f32::INFINITY, |a, &b| a.min(b));
613 let max = data.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
614 (min, max)
615 }
616 };
617
618 let params =
620 QuantizationParams::from_range(min_val, max_val, self.config.qtype, self.config.scheme);
621
622 let quantized_data = params.quantize_array(data);
624
625 Ok(QuantizedTensor::new(quantized_data, params, shape).with_name(name.unwrap_or_default()))
626 }
627
628 pub fn config(&self) -> &QuantizationConfig {
630 &self.config
631 }
632}
633
634#[derive(Debug, Clone, Serialize, Deserialize)]
636pub struct CompressionStats {
637 pub original_size: usize,
639 pub quantized_size: usize,
641 pub compression_ratio: f32,
643 pub num_quantized_layers: usize,
645 pub num_fp32_layers: usize,
647 pub avg_sqnr_db: f32,
649}
650
651impl CompressionStats {
652 pub fn new() -> Self {
654 Self {
655 original_size: 0,
656 quantized_size: 0,
657 compression_ratio: 1.0,
658 num_quantized_layers: 0,
659 num_fp32_layers: 0,
660 avg_sqnr_db: 0.0,
661 }
662 }
663
664 pub fn size_reduction_percent(&self) -> f32 {
666 (1.0 - (self.quantized_size as f32 / self.original_size as f32)) * 100.0
667 }
668}
669
670impl Default for CompressionStats {
671 fn default() -> Self {
672 Self::new()
673 }
674}
675
676#[cfg(test)]
677mod tests {
678 use super::*;
679
680 #[test]
681 fn test_quantization_type_bits() {
682 assert_eq!(QuantizationType::Int4.bits(), 4);
683 assert_eq!(QuantizationType::Int8.bits(), 8);
684 assert_eq!(QuantizationType::Fp16.bits(), 16);
685 }
686
687 #[test]
688 fn test_quantization_type_compression_ratio() {
689 assert_eq!(QuantizationType::Int8.compression_ratio(), 4.0); assert_eq!(QuantizationType::Int4.compression_ratio(), 8.0); assert_eq!(QuantizationType::Fp16.compression_ratio(), 2.0); }
693
694 #[test]
695 fn test_quantization_params_symmetric() {
696 let params = QuantizationParams::from_range(
697 -10.0,
698 10.0,
699 QuantizationType::Int8,
700 QuantizationScheme::Symmetric,
701 );
702
703 assert_eq!(params.zero_point, 0);
704 assert!(params.scale > 0.0);
705
706 let val = 5.0_f32;
708 let quantized = params.quantize(val);
709 let dequantized = params.dequantize(quantized);
710 assert!((val - dequantized).abs() < 0.1); }
712
713 #[test]
714 fn test_quantization_params_asymmetric() {
715 let params = QuantizationParams::from_range(
716 0.0,
717 10.0,
718 QuantizationType::Uint8,
719 QuantizationScheme::Asymmetric,
720 );
721
722 assert!(params.scale > 0.0);
723
724 let val = 5.0_f32;
726 let quantized = params.quantize(val);
727 let dequantized = params.dequantize(quantized);
728 assert!((val - dequantized).abs() < 0.1);
729 }
730
731 #[test]
732 fn test_quantization_array() {
733 let params = QuantizationParams::from_range(
734 -1.0,
735 1.0,
736 QuantizationType::Int8,
737 QuantizationScheme::Symmetric,
738 );
739
740 let data = vec![-0.5, 0.0, 0.5, 1.0];
741 let quantized = params.quantize_array(&data);
742 let dequantized = params.dequantize_array(&quantized);
743
744 for (orig, deq) in data.iter().zip(dequantized.iter()) {
745 assert!((orig - deq).abs() < 0.1);
746 }
747 }
748
749 #[test]
750 fn test_quantized_tensor_creation() {
751 let data = vec![1, 2, 3, 4, 5, 6];
752 let params = QuantizationParams::new(0.1, 0, QuantizationType::Int8);
753 let shape = vec![2, 3];
754
755 let qtensor = QuantizedTensor::new(data.clone(), params.clone(), shape.clone());
756
757 assert_eq!(qtensor.quantized_data, data);
758 assert_eq!(qtensor.shape, shape);
759 assert!(qtensor.compression_ratio() > 1.0);
760 }
761
762 #[test]
763 fn test_quantized_tensor_size() {
764 let data = vec![0; 1000]; let params = QuantizationParams::new(1.0, 0, QuantizationType::Int8);
766 let qtensor = QuantizedTensor::new(data, params, vec![1000]);
767
768 assert_eq!(qtensor.size_bytes(), 1000);
770
771 let params_int4 = QuantizationParams::new(1.0, 0, QuantizationType::Int4);
772 let qtensor_int4 = QuantizedTensor::new(vec![0; 1000], params_int4, vec![1000]);
773
774 assert_eq!(qtensor_int4.size_bytes(), 500);
776 }
777
778 #[test]
779 fn test_calibration_dataset() {
780 let mut dataset = CalibrationDataset::new(100);
781
782 dataset.add_sample(vec![1.0, 2.0, 3.0]);
783 dataset.add_sample(vec![4.0, 5.0, 6.0]);
784
785 assert_eq!(dataset.samples().len(), 2);
786
787 let stats = dataset.compute_statistics();
788 assert_eq!(stats.min_val, 1.0);
789 assert_eq!(stats.max_val, 6.0);
790 assert_eq!(stats.num_samples, 6);
791 }
792
793 #[test]
794 fn test_calibration_percentile() {
795 let mut dataset = CalibrationDataset::new(100);
796
797 for i in 1..=100 {
799 dataset.add_sample(vec![i as f32]);
800 }
801
802 let p50 = dataset.percentile(50.0);
803 assert!((p50 - 50.0).abs() < 5.0); let p99 = dataset.percentile(99.0);
806 assert!(p99 > 95.0);
807 }
808
809 #[test]
810 fn test_quantization_config() {
811 let config = QuantizationConfig::new(QuantizationType::Int8)
812 .with_granularity(QuantizationGranularity::PerChannel)
813 .with_scheme(QuantizationScheme::Symmetric)
814 .skip_layer("layer1".to_string())
815 .force_fp32_layer("output".to_string());
816
817 assert_eq!(config.qtype, QuantizationType::Int8);
818 assert!(config.quantize_weights);
819 assert!(config.skip_layers.contains(&"layer1".to_string()));
820 }
821
822 #[test]
823 fn test_quantizer_basic() {
824 let config = QuantizationConfig::new(QuantizationType::Int8);
825 let quantizer = Quantizer::new(config);
826
827 let data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
828 let shape = vec![5];
829
830 let result = quantizer.quantize_tensor(&data, shape, Some("test".to_string()));
831 assert!(result.is_ok());
832
833 let qtensor = result.unwrap();
834 assert_eq!(qtensor.shape, vec![5]);
835 assert_eq!(qtensor.quantized_data.len(), 5);
836 }
837
838 #[test]
839 fn test_quantizer_with_calibration() {
840 let mut calib_data = CalibrationDataset::new(10);
841 calib_data.add_sample(vec![0.0, 1.0, 2.0, 3.0]);
842 calib_data.add_sample(vec![0.5, 1.5, 2.5, 3.5]);
843
844 let config = QuantizationConfig::new(QuantizationType::Int8).with_calibration(
845 CalibrationMethod::Percentile {
846 lower: 1.0,
847 upper: 99.0,
848 },
849 );
850
851 let quantizer = Quantizer::new(config).with_calibration_data(calib_data);
852
853 let data = vec![1.0, 2.0, 3.0];
854 let result = quantizer.quantize_tensor(&data, vec![3], None);
855 assert!(result.is_ok());
856 }
857
858 #[test]
859 fn test_quantization_error_metrics() {
860 let original = vec![1.0, 2.0, 3.0, 4.0, 5.0];
861
862 let params = QuantizationParams::from_range(
863 0.0,
864 5.0,
865 QuantizationType::Int8,
866 QuantizationScheme::Asymmetric,
867 );
868
869 let quantized_data = params.quantize_array(&original);
870 let qtensor = QuantizedTensor::new(quantized_data, params, vec![5]);
871
872 let error = qtensor.quantization_error(&original);
873
874 assert!(error.mse >= 0.0);
875 assert!(error.rmse >= 0.0);
876 assert!(error.max_error >= 0.0);
877 assert!(error.sqnr_db > 0.0); }
879
880 #[test]
881 fn test_compression_stats() {
882 let mut stats = CompressionStats::new();
883 stats.original_size = 1000;
884 stats.quantized_size = 250;
885 stats.compression_ratio = 4.0;
886
887 assert_eq!(stats.size_reduction_percent(), 75.0);
888 }
889
890 #[test]
891 fn test_quantizer_skip_layer() {
892 let config =
893 QuantizationConfig::new(QuantizationType::Int8).skip_layer("skip_me".to_string());
894
895 let quantizer = Quantizer::new(config);
896
897 let data = vec![1.0, 2.0, 3.0];
898 let result = quantizer.quantize_tensor(&data, vec![3], Some("skip_me".to_string()));
899
900 assert!(result.is_err()); }
902}