//! Early stopping strategies with rtx-validation integration //! //! This module provides early stopping functionality for AutoML optimization, //! including cross-validation based stopping using rtx-validation's KFold. use crate::{AutoMLError, AutoMLResult}; use rtx_tensor::Tensor; use rtx_validation::cv::{CrossValidator, KFold}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::time::Instant; /// Early stopping controller for AutoML optimization pub struct EarlyStopping { criteria: Option, trials: Vec, cv_trials: Vec, start_time: Option, best_score: f64, best_trial_id: Option, trials_since_improvement: usize, cv_folds: usize, multi_objectives: Vec, pareto_front: Vec, adaptive_patience_enabled: bool, } impl EarlyStopping { /// Create a new EarlyStopping instance pub fn new() -> AutoMLResult { Ok(Self { criteria: None, trials: Vec::new(), cv_trials: Vec::new(), start_time: None, best_score: f64::NEG_INFINITY, best_trial_id: None, trials_since_improvement: 0, cv_folds: 5, multi_objectives: Vec::new(), pareto_front: Vec::new(), adaptive_patience_enabled: false, }) } /// Set the stopping criteria pub fn set_criteria(&mut self, criteria: StoppingCriteria) { self.criteria = Some(criteria); } /// Record a trial result pub fn record_trial(&mut self, trial: TrialResult) { let score = trial.score; let trial_id = trial.trial_id; self.trials.push(trial); // Check for improvement if let Some(ref criteria) = self.criteria { let improvement = score - self.best_score; if improvement > criteria.min_improvement { self.best_score = score; self.best_trial_id = Some(trial_id); self.trials_since_improvement = 0; } else { self.trials_since_improvement += 1; } } else if score > self.best_score { self.best_score = score; self.best_trial_id = Some(trial_id); self.trials_since_improvement = 0; } else { self.trials_since_improvement += 1; } } /// Check if optimization should stop pub fn should_stop(&self) -> bool { if let Some(ref criteria) = self.criteria { // Check patience if self.trials_since_improvement >= criteria.patience { return true; } // Check max trials if self.trials.len() >= criteria.max_trials { return true; } // Check target score if let Some(target) = criteria.target_score && self.best_score >= target { return true; } // Check time budget if let Some(start) = self.start_time && start.elapsed().as_secs() >= criteria.max_time_seconds { return true; } } false } /// Start the optimization timer pub fn start_timer(&mut self) { self.start_time = Some(Instant::now()); } /// Get remaining time in seconds pub fn get_remaining_time_seconds(&self) -> f64 { if let Some(ref criteria) = self.criteria { if let Some(start) = self.start_time { let elapsed = start.elapsed().as_secs_f64(); let budget = criteria.max_time_seconds as f64; return (budget - elapsed).max(0.0); } return criteria.max_time_seconds as f64; } 3600.0 // Default 1 hour } /// Get the optimization history pub fn get_optimization_history(&self) -> OptimizationHistory { let total_time = self .start_time .map_or(0.0, |s| s.elapsed().as_secs_f64()); OptimizationHistory { trials: self.trials.clone(), best_score: self.best_score, best_trial_id: self.best_trial_id, total_time_seconds: total_time, } } /// Get current patience (trials without improvement) pub fn get_current_patience(&self) -> usize { self.trials_since_improvement } /// Enable adaptive patience based on learning dynamics pub fn enable_adaptive_patience(&mut self, enabled: bool) { self.adaptive_patience_enabled = enabled; } /// Enable multi-objective optimization pub fn enable_multi_objective(&mut self, objectives: Vec) { self.multi_objectives = objectives; } /// Record a multi-objective trial pub fn record_multi_objective_trial( &mut self, trial_id: usize, objectives: HashMap, ) { // Create trial result from objectives let score = objectives.values().sum::() / objectives.len() as f64; let trial = TrialResult { trial_id, score, training_time_seconds: 0.0, memory_usage_mb: 0.0, hyperparameters: HashMap::new(), model_name: String::new(), }; // Update Pareto front self.update_pareto_front(trial.clone(), &objectives); self.trials.push(trial); } fn update_pareto_front(&mut self, trial: TrialResult, _objectives: &HashMap) { // Simplified Pareto dominance check // In a full implementation, this would check dominance across all objectives let dominated = self.pareto_front.iter().any(|p| p.score > trial.score); if !dominated { self.pareto_front.retain(|p| p.score >= trial.score); self.pareto_front.push(trial); } } /// Check if multi-objective optimization should stop pub fn should_stop_multi_objective(&self) -> AutoMLResult { // Stop if Pareto front hasn't changed recently if let Some(ref criteria) = self.criteria && self.trials_since_improvement >= criteria.patience { return Ok(true); } Ok(false) } /// Get the Pareto front of solutions pub fn get_pareto_front(&self) -> Vec { self.pareto_front.clone() } // ============================================================ // CV-BASED STOPPING WITH rtx-validation INTEGRATION // ============================================================ /// Enable CV-based stopping with KFold from rtx-validation pub fn enable_cv_based_stopping(&mut self, folds: usize) { self.cv_folds = folds; } /// Record a CV trial result with proper validation pub fn record_cv_trial(&mut self, trial_id: usize, cv_scores: Vec, mean: f64, std: f64) { let cv_result = CVTrialResult { trial_id, cv_scores, mean_score: mean, std_score: std, }; self.cv_trials.push(cv_result); // Update best score with CV confidence adjustment // Use lower bound estimate (mean - std) for conservative comparison let confidence_adjusted_score = mean - std; if confidence_adjusted_score > self.best_score { self.best_score = confidence_adjusted_score; self.best_trial_id = Some(trial_id); self.trials_since_improvement = 0; } else { self.trials_since_improvement += 1; } } /// Check if should stop based on CV confidence intervals pub fn should_stop_with_cv_confidence(&self) -> AutoMLResult { if let Some(ref criteria) = self.criteria { // Check patience with CV-adjusted scores if self.trials_since_improvement >= criteria.patience { return Ok(true); } // Check if recent trials show no statistically significant improvement if self.cv_trials.len() >= 3 { let recent_trials: Vec<_> = self.cv_trials.iter().rev().take(3).collect(); let all_overlap = recent_trials.windows(2).all(|pair| { let (t1, t2) = (&pair[0], &pair[1]); // Check if confidence intervals overlap (no significant improvement) let t1_upper = t1.mean_score + t1.std_score; let t2_lower = t2.mean_score - t2.std_score; t1_upper >= t2_lower }); if all_overlap && self.cv_trials.len() >= criteria.min_trials { return Ok(true); } } } Ok(false) } /// Get CV history for analysis pub fn get_cv_history(&self) -> Vec { self.cv_trials.clone() } /// Perform CV evaluation using rtx-validation KFold /// /// This method creates KFold splits and evaluates using the provided function. /// Returns (mean_score, std_score, fold_scores). pub fn evaluate_with_cv( &self, x: &Tensor, y: &Tensor, mut eval_fn: F, ) -> AutoMLResult<(f64, f64, Vec)> where F: FnMut(&[usize], &[usize]) -> AutoMLResult, { let cv = KFold::new(self.cv_folds).shuffle(true).random_state(42); let splits = cv .split(x, Some(y)) .map_err(|e| AutoMLError::ValidationError(e.to_string()))?; let mut fold_scores = Vec::new(); for (train_idx, val_idx) in splits { let score = eval_fn(&train_idx, &val_idx)?; fold_scores.push(score); } let mean = fold_scores.iter().sum::() / fold_scores.len() as f64; let variance = fold_scores.iter().map(|s| (s - mean).powi(2)).sum::() / fold_scores.len() as f64; let std = variance.sqrt(); Ok((mean, std, fold_scores)) } /// Serialize to JSON pub fn to_json(&self) -> AutoMLResult { let history = self.get_optimization_history(); serde_json::to_string(&history).map_err(|e| AutoMLError::SerializationError(e.to_string())) } /// Deserialize from JSON (creates new instance with history) pub fn from_json(_json: &str) -> AutoMLResult { // For now, return a fresh instance // Full implementation would restore state from JSON Self::new() } } impl Default for EarlyStopping { fn default() -> Self { Self::new().unwrap() } } /// Stopping criteria configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct StoppingCriteria { /// Number of trials without improvement before stopping pub patience: usize, /// Minimum improvement to reset patience pub min_improvement: f64, /// Maximum number of trials pub max_trials: usize, /// Maximum time in seconds pub max_time_seconds: u64, /// Target score to achieve (optional) pub target_score: Option, /// Minimum trials before stopping pub min_trials: usize, } impl Default for StoppingCriteria { fn default() -> Self { Self { patience: 10, min_improvement: 0.001, max_trials: 100, max_time_seconds: 3600, target_score: None, min_trials: 5, } } } /// Result of a single trial #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TrialResult { pub trial_id: usize, pub score: f64, pub training_time_seconds: f64, pub memory_usage_mb: f64, pub hyperparameters: HashMap, pub model_name: String, } /// Optimization history #[derive(Debug, Clone, Serialize, Deserialize)] pub struct OptimizationHistory { pub trials: Vec, pub best_score: f64, pub best_trial_id: Option, pub total_time_seconds: f64, } /// CV trial result with statistical information #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CVTrialResult { pub trial_id: usize, pub cv_scores: Vec, pub mean_score: f64, pub std_score: f64, } #[cfg(test)] mod tests { use super::*; #[test] fn test_early_stopping_creation() { let es = EarlyStopping::new().unwrap(); assert_eq!(es.get_current_patience(), 0); assert!(!es.should_stop()); } #[test] fn test_trial_recording() { let mut es = EarlyStopping::new().unwrap(); es.set_criteria(StoppingCriteria { patience: 3, min_improvement: 0.01, ..Default::default() }); // Record improving trials es.record_trial(TrialResult { trial_id: 0, score: 0.5, training_time_seconds: 1.0, memory_usage_mb: 100.0, hyperparameters: HashMap::new(), model_name: "test".to_string(), }); assert_eq!(es.best_score, 0.5); assert_eq!(es.get_current_patience(), 0); // Record non-improving trial es.record_trial(TrialResult { trial_id: 1, score: 0.5, training_time_seconds: 1.0, memory_usage_mb: 100.0, hyperparameters: HashMap::new(), model_name: "test".to_string(), }); assert_eq!(es.get_current_patience(), 1); assert!(!es.should_stop()); // Record more non-improving trials to trigger stopping es.record_trial(TrialResult { trial_id: 2, score: 0.5, ..Default::default() }); es.record_trial(TrialResult { trial_id: 3, score: 0.5, ..Default::default() }); assert!(es.should_stop()); } #[test] fn test_cv_based_stopping() { let mut es = EarlyStopping::new().unwrap(); es.enable_cv_based_stopping(5); es.set_criteria(StoppingCriteria { patience: 2, min_trials: 3, ..Default::default() }); // Record CV trials es.record_cv_trial(0, vec![0.8, 0.82, 0.79, 0.81, 0.80], 0.804, 0.011); assert!(!es.should_stop_with_cv_confidence().unwrap()); es.record_cv_trial(1, vec![0.81, 0.83, 0.80, 0.82, 0.81], 0.814, 0.012); assert!(!es.should_stop_with_cv_confidence().unwrap()); // Stagnating performance es.record_cv_trial(2, vec![0.81, 0.82, 0.80, 0.81, 0.80], 0.808, 0.008); es.record_cv_trial(3, vec![0.80, 0.81, 0.79, 0.80, 0.80], 0.800, 0.007); // Should detect overlapping confidence intervals let history = es.get_cv_history(); assert_eq!(history.len(), 4); } } impl Default for TrialResult { fn default() -> Self { Self { trial_id: 0, score: 0.0, training_time_seconds: 0.0, memory_usage_mb: 0.0, hyperparameters: HashMap::new(), model_name: String::new(), } } }