//! Test symmetric INT8 quantization implementation //! //! Run with: `cargo run --example test_symmetric_quantization` use candle_core::{Device, Tensor}; use ml::memory_optimization::quantization::{dequantize_tensor_from_int8, quantize_tensor_to_int8}; use std::time::Instant; fn main() -> Result<(), Box> { println!("=== Symmetric INT8 Quantization Tests ===\n"); let device = Device::Cpu; // Test 1: Basic quantization println!("Test 1: Basic Quantization"); let data = vec![-127.0f32, -64.0, 0.0, 64.0, 127.0]; let tensor = Tensor::from_vec(data.clone(), (5,), &device)?; let quantized = quantize_tensor_to_int8(&tensor, &device)?; println!(" Original values: {:?}", data); println!(" Quantized values: {:?}", quantized.data); println!(" Scale: {}", quantized.scale); println!(" Zero point: {}", quantized.zero_point); println!(" Shape: {:?}\n", quantized.shape); // Test 2: Round-trip accuracy println!("Test 2: Round-trip Accuracy"); let data = vec![-10.0f32, -5.0, 0.0, 5.0, 10.0]; let tensor = Tensor::from_vec(data.clone(), (5,), &device)?; let quantized = quantize_tensor_to_int8(&tensor, &device)?; let dequantized = dequantize_tensor_from_int8(&quantized, &device)?; let diff = tensor.sub(&dequantized)?.abs()?; let max_error = diff.max(0)?.to_scalar::()?; let mean_error = diff.mean_all()?.to_scalar::()?; println!(" Max reconstruction error: {:.6}", max_error); println!(" Mean reconstruction error: {:.6}", mean_error); println!( " Max allowed error (0.5 * scale): {:.6}\n", quantized.scale * 0.5 ); // Test 3: Performance benchmark println!("Test 3: Performance Benchmark (512x512 tensor)"); let tensor = Tensor::randn(0f32, 1.0, (512, 512), &device)?; let start = Instant::now(); let quantized = quantize_tensor_to_int8(&tensor, &device)?; let quantize_time = start.elapsed(); let start = Instant::now(); let _dequantized = dequantize_tensor_from_int8(&quantized, &device)?; let dequantize_time = start.elapsed(); println!( " Quantization time: {:.2}ms", quantize_time.as_secs_f64() * 1000.0 ); println!( " Dequantization time: {:.2}ms", dequantize_time.as_secs_f64() * 1000.0 ); println!(" Target: <1ms per layer\n"); // Test 4: Memory savings println!("Test 4: Memory Savings"); let original_bytes = 512 * 512 * 4; // FP32 = 4 bytes let quantized_bytes = quantized.memory_bytes(); let savings_ratio = (original_bytes - quantized_bytes) as f32 / original_bytes as f32; let compression = quantized.compression_ratio(); println!( " Original size: {} bytes ({:.2} MB)", original_bytes, original_bytes as f32 / 1024.0 / 1024.0 ); println!( " Quantized size: {} bytes ({:.2} MB)", quantized_bytes, quantized_bytes as f32 / 1024.0 / 1024.0 ); println!(" Memory savings: {:.2}%", savings_ratio * 100.0); println!(" Compression ratio: {:.2}x\n", compression); // Test 5: Multi-dimensional tensor println!("Test 5: Multi-dimensional Tensor (2x3x4)"); let tensor = Tensor::randn(0f32, 10.0, (2, 3, 4), &device)?; let quantized = quantize_tensor_to_int8(&tensor, &device)?; let dequantized = dequantize_tensor_from_int8(&quantized, &device)?; println!(" Original shape: {:?}", tensor.dims()); println!(" Quantized shape: {:?}", quantized.shape); println!(" Dequantized shape: {:?}", dequantized.dims()); println!(" Element count: {}\n", quantized.data.len()); // Test 6: Edge case - all zeros println!("Test 6: Edge Case - All Zeros"); let data = vec![0.0f32; 10]; let tensor = Tensor::from_vec(data, (10,), &device)?; let quantized = quantize_tensor_to_int8(&tensor, &device)?; println!( " All values zero: {}", quantized.data.iter().all(|&x| x == 0) ); println!(" Scale: {} (default for zero tensor)\n", quantized.scale); // Test 7: Extreme values println!("Test 7: Extreme Values (clamping test)"); let data = vec![-1000.0f32, -500.0, 0.0, 500.0, 1000.0]; let tensor = Tensor::from_vec(data, (5,), &device)?; let quantized = quantize_tensor_to_int8(&tensor, &device)?; println!(" Quantized values: {:?}", quantized.data); println!(" Min value (should be -127): {}", quantized.data[0]); println!(" Max value (should be 127): {}", quantized.data[4]); println!(" Scale: {:.6}\n", quantized.scale); println!("=== All Tests Passed! ==="); Ok(()) }