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

408 lines
13 KiB
Rust

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::<f32>()?;
// 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<String, u8>) -> HashMap<u8, usize> {
let mut counts = HashMap::new();
for precision in precisions.values() {
*counts.entry(*precision).or_insert(0) += 1;
}
counts
}
fn calculate_throughput(precisions: &HashMap<String, u8>) -> 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::<f64>()
/ precisions.len() as f64
}