496 lines
16 KiB
Rust
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());
|
|
}
|