use anyhow::Result; use rtx_compress::quantization::mixed_precision::{ LayerSensitivity, MixedPrecisionOptimizer, OptimizationObjective, PrecisionConfig, }; use rtx_tensor::{Device, Shape, Tensor}; use std::collections::HashMap; #[test] #[ignore = "Pre-existing Metal device randn issue"] fn test_layer_sensitivity_analysis() -> Result<()> { let device = Device::try_default()?; // Mock model layers let mut layers = HashMap::new(); layers.insert( "embeddings".to_string(), Tensor::randn(&[30000, 512], &device)?, ); layers.insert( "attention.query".to_string(), Tensor::randn(&[512, 512], &device)?, ); layers.insert( "attention.key".to_string(), Tensor::randn(&[512, 512], &device)?, ); layers.insert( "mlp.up_proj".to_string(), Tensor::randn(&[512, 2048], &device)?, ); layers.insert("layer_norm".to_string(), Tensor::randn(&[512], &device)?); let config = PrecisionConfig { precision_bits: vec![4, 8, 16], sensitivity_threshold: 0.02, performance_weight: 0.7, quality_weight: 0.3, }; let mut optimizer = MixedPrecisionOptimizer::new(config); // Mock calibration data for sensitivity analysis let calibration_data = Tensor::randn(&[100, 512], &device)?; // Analyze layer sensitivities let sensitivities = optimizer.analyze_sensitivity(&layers, &calibration_data)?; // Verify sensitivity analysis results assert_eq!(sensitivities.len(), layers.len()); // Layer norm should be most sensitive (normalization layers typically are) let layer_norm_sensitivity = sensitivities.get("layer_norm").unwrap(); assert!( layer_norm_sensitivity.error_impact > 0.05, "Layer norm should have high sensitivity" ); // Embeddings should be least sensitive (large, redundant parameters) let embeddings_sensitivity = sensitivities.get("embeddings").unwrap(); assert!( embeddings_sensitivity.error_impact < layer_norm_sensitivity.error_impact, "Embeddings should be less sensitive than layer norm" ); // Verify sensitivity ordering makes sense let mut sorted_sensitivities: Vec<_> = sensitivities.iter().collect(); sorted_sensitivities.sort_by(|a, b| b.1.error_impact.partial_cmp(&a.1.error_impact).unwrap()); println!("Layer sensitivity ranking:"); for (layer, sensitivity) in &sorted_sensitivities { println!(" {}: {:.4}", layer, sensitivity.error_impact); } Ok(()) } #[test] #[ignore = "Pre-existing Metal device randn issue"] fn test_precision_search_optimization() -> Result<()> { let device = Device::try_default()?; let mut layers = HashMap::new(); for i in 0..12 { // 12-layer transformer layers.insert( format!("layers.{}.attention", i), Tensor::randn(&[768, 768], &device)?, ); layers.insert( format!("layers.{}.mlp", i), Tensor::randn(&[768, 3072], &device)?, ); } let config = PrecisionConfig { precision_bits: vec![4, 6, 8, 12, 16], sensitivity_threshold: 0.01, performance_weight: 0.6, quality_weight: 0.4, }; let mut optimizer = MixedPrecisionOptimizer::new(config); let calibration_data = Tensor::randn(&[50, 768], &device)?; // Set optimization objective let objective = OptimizationObjective { target_compression_ratio: 4.0, max_quality_loss: 0.05, memory_constraint_mb: Some(500), }; // Run optimization let optimal_config = optimizer.optimize(&layers, &calibration_data, objective)?; // Verify optimization results assert_eq!(optimal_config.layer_precisions.len(), layers.len()); // Calculate achieved compression ratio let original_bits = layers.len() * 32; // All fp32 let compressed_bits: usize = optimal_config .layer_precisions .values() .map(|&b| b as usize) .sum(); let actual_ratio = original_bits as f32 / compressed_bits as f32; assert!( actual_ratio >= 3.5, "Should achieve at least 3.5x compression" ); assert!( optimal_config.estimated_quality_loss <= 0.06, "Quality loss should be within bounds" ); // Verify precision assignment makes sense // Earlier layers often need higher precision let layer_0_precision = optimal_config .layer_precisions .get("layers.0.attention") .unwrap(); let layer_11_precision = optimal_config .layer_precisions .get("layers.11.attention") .unwrap(); // This is a heuristic - may not always hold, but generally true println!( "Layer 0 precision: {}, Layer 11 precision: {}", layer_0_precision, layer_11_precision ); Ok(()) } #[test] #[ignore = "Pre-existing Metal device randn issue"] fn test_dynamic_precision_adjustment() -> Result<()> { let device = Device::try_default()?; let mut layers = HashMap::new(); layers.insert( "critical_layer".to_string(), Tensor::randn(&[256, 256], &device)?, ); layers.insert( "redundant_layer".to_string(), Tensor::randn(&[1024, 1024], &device)?, ); let config = PrecisionConfig { precision_bits: vec![4, 8, 12, 16], sensitivity_threshold: 0.015, performance_weight: 0.5, quality_weight: 0.5, }; let mut optimizer = MixedPrecisionOptimizer::new(config); // Initial calibration let calibration_data = Tensor::randn(&[100, 256], &device)?; let sensitivities = optimizer.analyze_sensitivity(&layers, &calibration_data)?; // Simulate runtime performance feedback let mut performance_feedback = HashMap::new(); performance_feedback.insert("critical_layer".to_string(), 0.95); // High accuracy needed performance_feedback.insert("redundant_layer".to_string(), 0.85); // Lower accuracy ok // Adjust precisions based on runtime feedback let adjusted_config = optimizer.adjust_precisions_runtime(&sensitivities, &performance_feedback)?; // Critical layer should get higher precision let critical_precision = adjusted_config .layer_precisions .get("critical_layer") .unwrap(); let redundant_precision = adjusted_config .layer_precisions .get("redundant_layer") .unwrap(); assert!( critical_precision >= redundant_precision, "Critical layer should have >= precision than redundant layer" ); // Verify precision bounds assert!( *critical_precision >= 8, "Critical layer needs at least 8-bit" ); assert!(*redundant_precision >= 4, "All layers need at least 4-bit"); Ok(()) } #[test] #[ignore = "Pre-existing Metal device randn issue"] fn test_quantization_aware_training_simulation() -> Result<()> { let device = Device::try_default()?; let mut layers = HashMap::new(); layers.insert("layer1".to_string(), Tensor::randn(&[128, 256], &device)?); layers.insert("layer2".to_string(), Tensor::randn(&[256, 512], &device)?); let config = PrecisionConfig { precision_bits: vec![6, 8, 12], sensitivity_threshold: 0.02, performance_weight: 0.4, quality_weight: 0.6, }; let mut optimizer = MixedPrecisionOptimizer::new(config); optimizer.enable_quantization_aware_mode(true); // Simulate training iterations with precision updates let mut current_precisions = HashMap::new(); current_precisions.insert("layer1".to_string(), 8); current_precisions.insert("layer2".to_string(), 8); for iteration in 0..10 { // Generate training batch let batch_data = Tensor::randn(&[32, 128], &device)?; // Simulate forward pass with current precisions let layer1_quantized = optimizer.simulate_quantization(&layers["layer1"], current_precisions["layer1"])?; // Compute simulated loss (simplified) let two = Tensor::from_data(vec![2.0], vec![1], &device)?; let loss = layer1_quantized .pow(&two)? .mean(&[], false)? .to_scalar::()?; // Update precisions based on loss gradients if iteration > 5 && loss > 0.1 { // Increase precision if loss is high let new_precision = (current_precisions["layer1"] + 2).min(12); current_precisions.insert("layer1".to_string(), new_precision); } } // Verify precision was adjusted during "training" let final_layer1_precision = current_precisions["layer1"]; println!("Final layer1 precision: {}", final_layer1_precision); // Should either stay at 8 or increase based on loss feedback assert!(final_layer1_precision >= 8 && final_layer1_precision <= 12); Ok(()) } #[test] #[ignore = "Pre-existing Metal device randn issue"] fn test_hardware_aware_precision_mapping() -> Result<()> { let device = Device::try_default()?; let mut layers = HashMap::new(); for i in 0..8 { layers.insert( format!("conv_{}", i), Tensor::randn(&[64, 128, 3, 3], &device)?, ); } let config = PrecisionConfig { precision_bits: vec![4, 8, 16], sensitivity_threshold: 0.01, performance_weight: 0.8, // Prioritize performance quality_weight: 0.2, }; let mut optimizer = MixedPrecisionOptimizer::new(config); // Configure hardware constraints optimizer.set_hardware_constraints(&[ ("int4_ops_per_sec", 1_000_000.0), ("int8_ops_per_sec", 500000.0), ("fp16_ops_per_sec", 200000.0), ])?; let calibration_data = Tensor::randn(&[100, 64], &device)?; let objective = OptimizationObjective { target_compression_ratio: 3.0, max_quality_loss: 0.08, memory_constraint_mb: None, }; let optimal_config = optimizer.optimize(&layers, &calibration_data, objective)?; // Verify hardware-aware decisions let precision_counts = count_precision_usage(&optimal_config.layer_precisions); // Should prefer 4-bit (fastest) when quality allows assert!( precision_counts.get(&4).unwrap_or(&0) > &0, "Should use some 4-bit quantization for performance" ); // Calculate theoretical throughput let estimated_ops_per_sec = calculate_throughput(&optimal_config.layer_precisions); println!("Estimated throughput: {:.0} ops/sec", estimated_ops_per_sec); assert!( estimated_ops_per_sec > 300000.0, "Should achieve reasonable throughput with hardware constraints" ); Ok(()) } #[test] #[ignore = "Pre-existing Metal device randn issue"] fn test_precision_config_serialization() -> Result<()> { let device = Device::try_default()?; let mut layers = HashMap::new(); layers.insert( "test_layer".to_string(), Tensor::randn(&[64, 128], &device)?, ); let config = PrecisionConfig { precision_bits: vec![4, 8], sensitivity_threshold: 0.03, performance_weight: 0.6, quality_weight: 0.4, }; let mut optimizer = MixedPrecisionOptimizer::new(config); let calibration_data = Tensor::randn(&[50, 64], &device)?; let sensitivities = optimizer.analyze_sensitivity(&layers, &calibration_data)?; // Serialize sensitivity analysis results let serialized = optimizer.serialize_sensitivities(&sensitivities)?; // Create new optimizer and deserialize let mut optimizer2 = MixedPrecisionOptimizer::new(PrecisionConfig { precision_bits: vec![4, 8], sensitivity_threshold: 0.03, performance_weight: 0.6, quality_weight: 0.4, }); let deserialized_sensitivities = optimizer2.deserialize_sensitivities(&serialized)?; // Verify serialization preserved data assert_eq!(deserialized_sensitivities.len(), sensitivities.len()); for (layer_name, original_sens) in &sensitivities { let deserialized_sens = deserialized_sensitivities.get(layer_name).unwrap(); let error_diff = (original_sens.error_impact - deserialized_sens.error_impact).abs(); assert!( error_diff < 1e-6, "Sensitivity data should be preserved exactly" ); } Ok(()) } fn count_precision_usage(precisions: &HashMap) -> HashMap { let mut counts = HashMap::new(); for precision in precisions.values() { *counts.entry(*precision).or_insert(0) += 1; } counts } fn calculate_throughput(precisions: &HashMap) -> f64 { // Simplified throughput calculation based on precision precisions .values() .map(|&bits| match bits { 4 => 1_000_000.0, 8 => 500000.0, 16 => 200000.0, _ => 100000.0, }) .sum::() / precisions.len() as f64 }