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

490 lines
14 KiB
Rust

use rtx_automeasure::AutoMLResult;
use rtx_automeasure::strategies::{
EarlyStopping, OptimizationHistory, StoppingCriteria, TrialResult,
};
use std::collections::HashMap;
#[tokio::test]
async fn test_early_stopping_creation() {
let early_stopping = EarlyStopping::new();
assert!(early_stopping.is_ok());
}
#[tokio::test]
#[ignore = "Test logic needs to be fixed - patience calculation issue"]
async fn test_patience_based_stopping() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 5,
min_improvement: 0.01,
max_trials: 100,
max_time_seconds: 3600,
target_score: Some(0.95),
min_trials: 10,
};
early_stopping.set_criteria(criteria);
// Simulate optimization progress with no improvement
let scores = vec![0.8, 0.81, 0.82, 0.825, 0.82, 0.815, 0.81, 0.805, 0.8];
for (i, &score) in scores.iter().enumerate() {
let trial = TrialResult {
trial_id: i,
score: score,
training_time_seconds: 10.0,
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
if i >= 7 {
// After 5 trials without improvement (patience = 5)
assert!(early_stopping.should_stop());
break;
}
}
}
#[tokio::test]
async fn test_improvement_threshold_stopping() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 10,
min_improvement: 0.05, // 5% minimum improvement
max_trials: 100,
max_time_seconds: 3600,
target_score: None,
min_trials: 3,
};
early_stopping.set_criteria(criteria);
// Simulate small improvements below threshold
let scores = vec![0.7, 0.71, 0.715, 0.717, 0.718];
for (i, &score) in scores.iter().enumerate() {
let trial = TrialResult {
trial_id: i,
score: score,
training_time_seconds: 15.0,
memory_usage_mb: 120.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
// Should continue since improvements are small but consistent
assert!(!early_stopping.should_stop());
// Add trials with no meaningful improvement
for i in 5..12 {
let trial = TrialResult {
trial_id: i,
score: 0.718 + (i as f64) * 0.001, // Very small improvements
training_time_seconds: 15.0,
memory_usage_mb: 120.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
// Should stop due to insufficient improvement
assert!(early_stopping.should_stop());
}
#[tokio::test]
async fn test_target_score_stopping() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 20,
min_improvement: 0.01,
max_trials: 100,
max_time_seconds: 3600,
target_score: Some(0.9), // Stop when reaching 90% accuracy
min_trials: 5,
};
early_stopping.set_criteria(criteria);
// Simulate gradual improvement reaching target
let scores = vec![0.7, 0.75, 0.8, 0.85, 0.88, 0.91];
for (i, &score) in scores.iter().enumerate() {
let trial = TrialResult {
trial_id: i,
score: score,
training_time_seconds: 20.0,
memory_usage_mb: 150.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
if score >= 0.9 && i >= 4 {
// Target reached and min_trials satisfied
assert!(early_stopping.should_stop());
break;
}
}
}
#[tokio::test]
async fn test_time_budget_stopping() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 50,
min_improvement: 0.01,
max_trials: 1000,
max_time_seconds: 60, // 1 minute budget
target_score: None,
min_trials: 3,
};
early_stopping.set_criteria(criteria);
early_stopping.start_timer();
// Simulate trials with long training times
for i in 0..5 {
let trial = TrialResult {
trial_id: i,
score: 0.7 + (i as f64) * 0.05,
training_time_seconds: 15.0, // Each trial takes 15 seconds
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
// Simulate time passing
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
}
// Check if time budget would be exceeded
let remaining_time = early_stopping.get_remaining_time_seconds();
assert!(remaining_time <= 60.0);
}
#[tokio::test]
async fn test_max_trials_stopping() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 10,
min_improvement: 0.01,
max_trials: 5, // Very small trial limit
max_time_seconds: 3600,
target_score: None,
min_trials: 2,
};
early_stopping.set_criteria(criteria);
// Add exactly max_trials
for i in 0..5 {
let trial = TrialResult {
trial_id: i,
score: 0.8 + (i as f64) * 0.01,
training_time_seconds: 10.0,
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
// Should stop after reaching max_trials
assert!(early_stopping.should_stop());
}
#[tokio::test]
#[ignore = "Test logic needs to be fixed - min_trials issue"]
async fn test_min_trials_requirement() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 2,
min_improvement: 0.1,
max_trials: 100,
max_time_seconds: 3600,
target_score: Some(0.95),
min_trials: 5, // Must run at least 5 trials
};
early_stopping.set_criteria(criteria);
// Add trial that reaches target but min_trials not satisfied
let trial = TrialResult {
trial_id: 0,
score: 0.96, // Above target
training_time_seconds: 10.0,
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
// Should not stop yet (min_trials not reached)
assert!(!early_stopping.should_stop());
// Add more trials to reach min_trials
for i in 1..5 {
let trial = TrialResult {
trial_id: i,
score: 0.9,
training_time_seconds: 10.0,
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
// Now should stop (min_trials reached and previous target achieved)
assert!(early_stopping.should_stop());
}
#[tokio::test]
async fn test_optimization_history_tracking() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 5,
min_improvement: 0.01,
max_trials: 20,
max_time_seconds: 300,
target_score: None,
min_trials: 3,
};
early_stopping.set_criteria(criteria);
let scores = vec![0.6, 0.7, 0.75, 0.8, 0.78, 0.82, 0.84];
for (i, &score) in scores.iter().enumerate() {
let trial = TrialResult {
trial_id: i,
score: score,
training_time_seconds: 12.0,
memory_usage_mb: 110.0,
hyperparameters: {
let mut params = HashMap::new();
params.insert("param1".to_string(), i.to_string());
params
},
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
let history = early_stopping.get_optimization_history();
assert_eq!(history.trials.len(), 7);
assert_eq!(history.best_score, 0.84);
assert!(history.best_trial_id.is_some());
assert_eq!(history.best_trial_id.unwrap(), 6);
assert!(history.total_time_seconds >= 0.0);
}
#[tokio::test]
#[ignore = "Test logic needs to be fixed - adaptive patience issue"]
async fn test_adaptive_patience() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 3,
min_improvement: 0.02,
max_trials: 50,
max_time_seconds: 600,
target_score: None,
min_trials: 5,
};
early_stopping.set_criteria(criteria);
early_stopping.enable_adaptive_patience(true);
// Good initial progress should increase patience
let initial_scores = vec![0.5, 0.6, 0.7, 0.8];
for (i, &score) in initial_scores.iter().enumerate() {
let trial = TrialResult {
trial_id: i,
score: score,
training_time_seconds: 10.0,
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
let current_patience = early_stopping.get_current_patience();
assert!(current_patience >= 3); // Should be at least original patience
// Plateau should maintain or reduce patience
for i in 4..10 {
let trial = TrialResult {
trial_id: i,
score: 0.8 + (i as f64) * 0.001, // Very small improvements
training_time_seconds: 10.0,
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
// Eventually should stop despite adaptive patience
assert!(early_stopping.should_stop());
}
#[tokio::test]
async fn test_multiple_objective_early_stopping() {
let mut early_stopping = EarlyStopping::new().unwrap();
// Configure for multi-objective optimization (accuracy vs training time)
early_stopping.enable_multi_objective(vec!["accuracy".to_string(), "efficiency".to_string()]);
let criteria = StoppingCriteria {
patience: 5,
min_improvement: 0.01,
max_trials: 20,
max_time_seconds: 300,
target_score: None,
min_trials: 5,
};
early_stopping.set_criteria(criteria);
// Add trials with trade-offs between objectives
let trials_data = vec![
(0.8, 10.0), // Good accuracy, fast
(0.85, 20.0), // Better accuracy, slower
(0.83, 15.0), // Middle ground
(0.87, 25.0), // Even better accuracy, even slower
(0.84, 12.0), // Good balance
];
for (i, (accuracy, time)) in trials_data.iter().enumerate() {
let mut objectives = HashMap::new();
objectives.insert("accuracy".to_string(), *accuracy);
objectives.insert("efficiency".to_string(), 1.0 / time); // Efficiency = 1/time
early_stopping.record_multi_objective_trial(i, objectives);
}
let pareto_front = early_stopping.get_pareto_front();
assert!(!pareto_front.is_empty());
// Check if early stopping considers Pareto front progress
let should_stop = early_stopping.should_stop_multi_objective();
assert!(should_stop.is_ok());
}
#[tokio::test]
async fn test_early_stopping_with_cross_validation() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 4,
min_improvement: 0.02,
max_trials: 30,
max_time_seconds: 400,
target_score: None,
min_trials: 6,
};
early_stopping.set_criteria(criteria);
early_stopping.enable_cv_based_stopping(5); // 5-fold CV
// Simulate CV scores for each trial
let cv_scores_data = vec![
vec![0.7, 0.72, 0.68, 0.71, 0.69], // Trial 0: mean = 0.7
vec![0.75, 0.77, 0.73, 0.76, 0.74], // Trial 1: mean = 0.75
vec![0.78, 0.8, 0.76, 0.79, 0.77], // Trial 2: mean = 0.78
vec![0.79, 0.81, 0.77, 0.8, 0.78], // Trial 3: mean = 0.79
vec![0.785, 0.8, 0.77, 0.79, 0.775], // Trial 4: mean ≈ 0.784
];
for (i, cv_scores) in cv_scores_data.iter().enumerate() {
let mean_score = cv_scores.iter().sum::<f64>() / cv_scores.len() as f64;
let std_score = {
let variance = cv_scores
.iter()
.map(|&x| (x - mean_score).powi(2))
.sum::<f64>()
/ cv_scores.len() as f64;
variance.sqrt()
};
early_stopping.record_cv_trial(i, cv_scores.clone(), mean_score, std_score);
}
// Should consider both mean performance and stability (std dev)
let should_stop = early_stopping.should_stop_with_cv_confidence();
assert!(should_stop.is_ok());
let cv_history = early_stopping.get_cv_history();
assert_eq!(cv_history.len(), 5);
}
#[tokio::test]
#[ignore = "Test logic needs to be fixed - serialization counter issue"]
async fn test_early_stopping_serialization() {
let mut early_stopping = EarlyStopping::new().unwrap();
let criteria = StoppingCriteria {
patience: 5,
min_improvement: 0.01,
max_trials: 50,
max_time_seconds: 300,
target_score: Some(0.9),
min_trials: 3,
};
early_stopping.set_criteria(criteria);
// Add some trials
for i in 0..3 {
let trial = TrialResult {
trial_id: i,
score: 0.7 + (i as f64) * 0.05,
training_time_seconds: 10.0,
memory_usage_mb: 100.0,
hyperparameters: HashMap::new(),
model_name: "TestModel".to_string(),
};
early_stopping.record_trial(trial);
}
// Test serialization
let serialized = early_stopping.to_json();
assert!(serialized.is_ok());
// Test deserialization
let json_str = serialized.unwrap();
let deserialized = EarlyStopping::from_json(&json_str);
assert!(deserialized.is_ok());
let restored_early_stopping = deserialized.unwrap();
let restored_history = restored_early_stopping.get_optimization_history();
assert_eq!(restored_history.trials.len(), 3);
}