Files
rustytorch/crates/training/rtx-compress/src/quantization/tests.rs
T
2026-03-04 00:08:42 +00:00

496 lines
16 KiB
Rust

//! 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<f32> = (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<f32> = (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<Tensor> = (0..3)
.map(|_| {
let data: Vec<f32> = (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<f32> = (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<f32> = (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<f32> = (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<f32> = (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<f32> = (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<f32> = (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<f32> = (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<f32> = (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<f32> = (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());
}