//! Tests for post-training quantization use super::post_training::*; use rtx_tensor::{Device, Tensor}; #[test] fn test_quantization_scheme_bit_width() { assert_eq!(QuantizationScheme::INT8.bit_width(), 8); assert_eq!(QuantizationScheme::INT4.bit_width(), 4); assert_eq!(QuantizationScheme::INT2.bit_width(), 2); assert_eq!(QuantizationScheme::INT1.bit_width(), 1); } #[test] fn test_quantization_scheme_range() { assert_eq!(QuantizationScheme::INT8.range(), (-128, 127)); assert_eq!(QuantizationScheme::INT4.range(), (-8, 7)); assert_eq!(QuantizationScheme::INT2.range(), (-2, 1)); assert_eq!(QuantizationScheme::INT1.range(), (-1, 1)); } #[test] fn test_quantization_scheme_num_levels() { assert_eq!(QuantizationScheme::INT8.num_levels(), 256); assert_eq!(QuantizationScheme::INT4.num_levels(), 16); assert_eq!(QuantizationScheme::INT2.num_levels(), 4); assert_eq!(QuantizationScheme::INT1.num_levels(), 2); } #[test] fn test_calibration_config_new() { let config = CalibrationConfig::new(100, 0.99); assert_eq!(config.num_samples, 100); assert_eq!(config.coverage, 0.99); assert_eq!(config.method, CalibrationMethod::MinMax); } #[test] fn test_calibration_config_percentile() { let config = CalibrationConfig::new_with_percentile(100, 0.01, 0.99); assert_eq!(config.num_samples, 100); assert_eq!(config.method, CalibrationMethod::Percentile); assert_eq!(config.lower_percentile, 0.01); assert_eq!(config.upper_percentile, 0.99); } #[test] fn test_calibration_config_entropy() { let config = CalibrationConfig::new_entropy_based(100, 1024); assert_eq!(config.num_samples, 100); assert_eq!(config.method, CalibrationMethod::Entropy); assert_eq!(config.histogram_bins, 1024); } #[test] fn test_quantization_config_new() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); assert_eq!(config.scheme, QuantizationScheme::INT8); assert!(!config.symmetric); assert!(!config.per_channel); } #[test] fn test_quantization_config_symmetric() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new_symmetric(QuantizationScheme::INT8, calibration); assert!(config.symmetric); } #[test] fn test_quantization_config_per_channel() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new_per_channel(QuantizationScheme::INT8, calibration, 0); assert!(config.per_channel); assert_eq!(config.channel_dim, 0); } #[test] fn test_quantization_statistics_default() { let stats = QuantizationStatistics::default(); assert_eq!(stats.scale, 1.0); assert_eq!(stats.zero_point, 0); assert_eq!(stats.bit_width, 8); assert_eq!(stats.quantization_type, "per_tensor"); } #[test] fn test_quantized_tensor_storage_size() { let device = Device::try_default().unwrap(); let values = Tensor::zeros(&[256], &device).unwrap(); let qt_int8 = QuantizedTensor { values: values.clone(), scale: 1.0, zero_point: 0, shape: vec![256], bit_width: 8, per_channel_scales: None, per_channel_zero_points: None, }; assert_eq!(qt_int8.storage_size(), 256); let qt_int4 = QuantizedTensor { values: values.clone(), scale: 1.0, zero_point: 0, shape: vec![256], bit_width: 4, per_channel_scales: None, per_channel_zero_points: None, }; assert_eq!(qt_int4.storage_size(), 128); let qt_int2 = QuantizedTensor { values: values.clone(), scale: 1.0, zero_point: 0, shape: vec![256], bit_width: 2, per_channel_scales: None, per_channel_zero_points: None, }; assert_eq!(qt_int2.storage_size(), 64); let qt_int1 = QuantizedTensor { values, scale: 1.0, zero_point: 0, shape: vec![256], bit_width: 1, per_channel_scales: None, per_channel_zero_points: None, }; assert_eq!(qt_int1.storage_size(), 32); } #[test] fn test_quantized_tensor_compression_ratio() { let device = Device::try_default().unwrap(); let values = Tensor::zeros(&[256], &device).unwrap(); let qt_int8 = QuantizedTensor { values: values.clone(), scale: 1.0, zero_point: 0, shape: vec![256], bit_width: 8, per_channel_scales: None, per_channel_zero_points: None, }; assert_eq!(qt_int8.compression_ratio(), 4.0); // f32 (32-bit) / int8 (8-bit) let qt_int4 = QuantizedTensor { values, scale: 1.0, zero_point: 0, shape: vec![256], bit_width: 4, per_channel_scales: None, per_channel_zero_points: None, }; assert_eq!(qt_int4.compression_ratio(), 8.0); // f32 (32-bit) / int4 (4-bit) } #[test] fn test_quantized_tensor_is_per_channel() { let device = Device::try_default().unwrap(); let values = Tensor::zeros(&[256], &device).unwrap(); let qt_per_tensor = QuantizedTensor { values: values.clone(), scale: 1.0, zero_point: 0, shape: vec![256], bit_width: 8, per_channel_scales: None, per_channel_zero_points: None, }; assert!(!qt_per_tensor.is_per_channel()); let qt_per_channel = QuantizedTensor { values, scale: 0.0, zero_point: 0, shape: vec![4, 64], bit_width: 8, per_channel_scales: Some(vec![1.0, 1.0, 1.0, 1.0]), per_channel_zero_points: Some(vec![0, 0, 0, 0]), }; assert!(qt_per_channel.is_per_channel()); } #[test] fn test_post_training_quantizer_creation() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let quantizer = PostTrainingQuantizer::new(config); assert!(quantizer.is_ok()); } #[test] fn test_quantizer_calibrate_and_quantize() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); // Create calibration data let data: Vec = (0..256).map(|i| (i as f32 - 128.0) / 128.0).collect(); let calibration_tensor = Tensor::from_slice(&data, &[256], &device).unwrap(); // Calibrate let result = quantizer.calibrate(&[calibration_tensor.clone()]); assert!(result.is_ok()); // Quantize let quantized = quantizer.quantize(&calibration_tensor); assert!(quantized.is_ok()); let qt = quantized.unwrap(); assert_eq!(qt.bit_width, 8); assert_eq!(qt.shape, vec![256]); } #[test] fn test_quantizer_dequantize() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); // Create data with known values let data: Vec = (0..64).map(|i| i as f32 / 64.0).collect(); let original = Tensor::from_slice(&data, &[64], &device).unwrap(); // Calibrate and quantize quantizer.calibrate(&[original.clone()]).unwrap(); let quantized = quantizer.quantize(&original).unwrap(); // Dequantize let dequantized = quantizer.dequantize(&quantized); assert!(dequantized.is_ok()); let dequant_vals = dequantized.unwrap().to_vec().unwrap(); assert_eq!(dequant_vals.len(), 64); } #[test] fn test_quantizer_batch_operations() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); // Create batch of tensors let batch: Vec = (0..3) .map(|_| { let data: Vec = (0..64).map(|i| i as f32 / 64.0).collect(); Tensor::from_slice(&data, &[64], &device).unwrap() }) .collect(); // Calibrate quantizer.calibrate(&batch).unwrap(); // Quantize batch let quantized_batch = quantizer.quantize_batch(&batch); assert!(quantized_batch.is_ok()); assert_eq!(quantized_batch.unwrap().len(), 3); } #[test] fn test_quantizer_get_statistics() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); let data: Vec = (0..256).map(|i| (i as f32 - 128.0) / 128.0).collect(); let calibration_tensor = Tensor::from_slice(&data, &[256], &device).unwrap(); quantizer.calibrate(&[calibration_tensor]).unwrap(); let stats = quantizer.get_statistics(); assert_eq!(stats.bit_width, 8); assert_eq!(stats.calibration_samples, 1); assert_eq!(stats.quantization_type, "per_tensor"); } #[test] fn test_quantizer_pack_unpack_int8() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); let data: Vec = (0..64).map(|i| i as f32).collect(); let tensor = Tensor::from_slice(&data, &[64], &device).unwrap(); quantizer.calibrate(&[tensor.clone()]).unwrap(); let quantized = quantizer.quantize(&tensor).unwrap(); // Pack let packed = quantizer.pack_quantized_tensor(&quantized); assert!(packed.is_ok()); assert_eq!(packed.unwrap().len(), 64); // 64 bytes for INT8 } #[test] fn test_quantizer_pack_int4() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT4, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); let data: Vec = (0..64).map(|i| (i % 8) as f32 - 4.0).collect(); let tensor = Tensor::from_slice(&data, &[64], &device).unwrap(); quantizer.calibrate(&[tensor.clone()]).unwrap(); let quantized = quantizer.quantize(&tensor).unwrap(); // Pack let packed = quantizer.pack_quantized_tensor(&quantized); assert!(packed.is_ok()); assert_eq!(packed.unwrap().len(), 32); // 32 bytes for INT4 (2 values per byte) } #[test] fn test_quantization_error_metrics() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); let data: Vec = (0..64).map(|i| i as f32 / 64.0).collect(); let tensor = Tensor::from_slice(&data, &[64], &device).unwrap(); quantizer.calibrate(&[tensor.clone()]).unwrap(); let quantized = quantizer.quantize(&tensor).unwrap(); let metrics = quantizer.calculate_quantization_error(&tensor, &quantized); assert!(metrics.is_ok()); let m = metrics.unwrap(); assert!(m.mse >= 0.0); assert!(m.mae >= 0.0); assert!(m.max_error >= 0.0); assert_eq!(m.bit_width, 8); } #[test] fn test_quantizer_serialization() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); let data: Vec = (0..64).map(|i| i as f32 / 64.0).collect(); let tensor = Tensor::from_slice(&data, &[64], &device).unwrap(); quantizer.calibrate(&[tensor]).unwrap(); // Serialize let serialized = quantizer.serialize(); assert!(serialized.is_ok()); // Deserialize let deserialized = PostTrainingQuantizer::deserialize(&serialized.unwrap()); assert!(deserialized.is_ok()); let restored = deserialized.unwrap(); let stats = restored.get_statistics(); assert_eq!(stats.bit_width, 8); } #[test] fn test_symmetric_quantization() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new_symmetric(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); // Symmetric data around zero let data: Vec = (0..256).map(|i| (i as f32 - 128.0) / 128.0).collect(); let tensor = Tensor::from_slice(&data, &[256], &device).unwrap(); quantizer.calibrate(&[tensor.clone()]).unwrap(); let quantized = quantizer.quantize(&tensor).unwrap(); // For symmetric quantization, zero_point should be 0 assert_eq!(quantized.zero_point, 0); } #[test] fn test_percentile_calibration() { let calibration = CalibrationConfig::new_with_percentile(100, 0.01, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); // Data with outliers let mut data: Vec = (0..256).map(|i| i as f32 / 256.0).collect(); data[0] = -100.0; // Outlier data[255] = 100.0; // Outlier let tensor = Tensor::from_slice(&data, &[256], &device).unwrap(); let result = quantizer.calibrate(&[tensor]); assert!(result.is_ok()); let stats = quantizer.get_statistics(); // Percentile calibration should clip outliers assert!(stats.min_value > -100.0); assert!(stats.max_value < 100.0); } #[test] #[ignore = "Pre-existing entropy calibration issue"] fn test_entropy_calibration() { let calibration = CalibrationConfig::new_entropy_based(100, 256); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); let data: Vec = (0..256).map(|i| i as f32 / 256.0).collect(); let tensor = Tensor::from_slice(&data, &[256], &device).unwrap(); let result = quantizer.calibrate(&[tensor]); assert!(result.is_ok()); let stats = quantizer.get_statistics(); assert!(stats.entropy_score >= 0.0); } #[test] fn test_mixed_precision_config() { use std::collections::HashMap; let mut layer_configs = HashMap::new(); layer_configs.insert("layer1".to_string(), QuantizationScheme::INT8); layer_configs.insert("layer2".to_string(), QuantizationScheme::INT4); let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new_mixed_precision(layer_configs, calibration); assert!(config.mixed_precision.is_some()); let mp = config.mixed_precision.unwrap(); assert_eq!(mp.get("layer1"), Some(&QuantizationScheme::INT8)); assert_eq!(mp.get("layer2"), Some(&QuantizationScheme::INT4)); } #[test] fn test_quantizer_not_calibrated_error() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let quantizer = PostTrainingQuantizer::new(config).unwrap(); let device = Device::try_default().unwrap(); let data: Vec = (0..64).map(|i| i as f32).collect(); let tensor = Tensor::from_slice(&data, &[64], &device).unwrap(); // Should fail because quantizer is not calibrated let result = quantizer.quantize(&tensor); assert!(result.is_err()); } #[test] fn test_empty_calibration_data_error() { let calibration = CalibrationConfig::new(100, 0.99); let config = QuantizationConfig::new(QuantizationScheme::INT8, calibration); let mut quantizer = PostTrainingQuantizer::new(config).unwrap(); // Should fail with empty calibration data let result = quantizer.calibrate(&[]); assert!(result.is_err()); }