Skip to main content

trustformers_debug/model_diagnostics/
training.rs

1//! Training dynamics and convergence analysis.
2//!
3//! This module provides comprehensive training dynamics analysis including
4//! convergence detection, overfitting/underfitting identification, plateau
5//! detection, and training stability assessment for optimizing training processes.
6
7use std::collections::VecDeque;
8
9use super::types::{
10    ConvergenceStatus, ModelPerformanceMetrics, OverfittingIndicator, PlateauInfo,
11    TrainingDynamics, TrainingStability, UnderfittingIndicator,
12};
13
14/// Training dynamics analyzer for monitoring and analyzing training behavior.
15#[derive(Debug)]
16pub struct TrainingDynamicsAnalyzer {
17    /// Historical metrics for analysis
18    metrics_history: VecDeque<ModelPerformanceMetrics>,
19    /// Configuration for analysis thresholds
20    config: TrainingAnalysisConfig,
21    /// Current training state
22    current_state: TrainingState,
23}
24
25/// Configuration for training analysis.
26#[derive(Debug, Clone)]
27pub struct TrainingAnalysisConfig {
28    /// Window size for convergence analysis
29    pub convergence_window: usize,
30    /// Minimum improvement threshold for convergence
31    pub min_improvement_threshold: f64,
32    /// Maximum variance threshold for stability
33    pub max_variance_threshold: f64,
34    /// Minimum plateau duration to consider
35    pub min_plateau_duration: usize,
36    /// Train-validation gap threshold for overfitting
37    pub overfitting_gap_threshold: f64,
38    /// Minimum learning rate for underfitting detection
39    pub min_learning_rate: f64,
40}
41
42impl Default for TrainingAnalysisConfig {
43    fn default() -> Self {
44        Self {
45            convergence_window: 20,
46            min_improvement_threshold: 0.001,
47            max_variance_threshold: 0.1,
48            min_plateau_duration: 10,
49            overfitting_gap_threshold: 0.05,
50            min_learning_rate: 1e-6,
51        }
52    }
53}
54
55/// Current training state information.
56#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
57pub struct TrainingState {
58    /// Steps since last improvement
59    steps_since_improvement: usize,
60    /// Best loss achieved so far
61    best_loss: f64,
62    /// Current plateau information
63    current_plateau: Option<PlateauInfo>,
64    /// Convergence status history
65    convergence_history: VecDeque<ConvergenceStatus>,
66}
67
68impl Default for TrainingState {
69    fn default() -> Self {
70        Self {
71            steps_since_improvement: 0,
72            best_loss: f64::INFINITY,
73            current_plateau: None,
74            convergence_history: VecDeque::new(),
75        }
76    }
77}
78
79impl TrainingDynamicsAnalyzer {
80    /// Create a new training dynamics analyzer.
81    pub fn new() -> Self {
82        Self {
83            metrics_history: VecDeque::new(),
84            config: TrainingAnalysisConfig::default(),
85            current_state: TrainingState::default(),
86        }
87    }
88
89    /// Create a new analyzer with custom configuration.
90    pub fn with_config(config: TrainingAnalysisConfig) -> Self {
91        Self {
92            metrics_history: VecDeque::new(),
93            config,
94            current_state: TrainingState::default(),
95        }
96    }
97
98    /// Add new training metrics for analysis.
99    pub fn add_metrics(&mut self, metrics: ModelPerformanceMetrics) {
100        // Update training state
101        if metrics.loss < self.current_state.best_loss {
102            self.current_state.best_loss = metrics.loss;
103            self.current_state.steps_since_improvement = 0;
104        } else {
105            self.current_state.steps_since_improvement += 1;
106        }
107
108        self.metrics_history.push_back(metrics);
109
110        // Maintain reasonable history size
111        if self.metrics_history.len() > 1000 {
112            self.metrics_history.pop_front();
113        }
114
115        // Update convergence history
116        let status = self.detect_convergence_status();
117        self.current_state.convergence_history.push_back(status);
118        if self.current_state.convergence_history.len() > 50 {
119            self.current_state.convergence_history.pop_front();
120        }
121    }
122
123    /// Record training dynamics information.
124    pub fn record_training_dynamics(&mut self, _dynamics: TrainingDynamics) {
125        // Training dynamics are computed via analysis rather than stored directly
126        // This method is provided for API compatibility
127    }
128
129    /// Analyze current training dynamics.
130    pub fn analyze_training_dynamics(&self) -> TrainingDynamics {
131        let convergence_status = self.detect_convergence_status();
132        let training_stability = self.assess_training_stability();
133        let learning_efficiency = self.calculate_learning_efficiency();
134        let overfitting_indicators = self.detect_overfitting_indicators();
135        let underfitting_indicators = self.detect_underfitting_indicators();
136        let plateau_detection = self.detect_plateau();
137
138        TrainingDynamics {
139            convergence_status,
140            training_stability,
141            learning_efficiency,
142            overfitting_indicators,
143            underfitting_indicators,
144            plateau_detection,
145        }
146    }
147
148    /// Detect current convergence status.
149    pub fn detect_convergence_status(&self) -> ConvergenceStatus {
150        if self.metrics_history.len() < self.config.convergence_window {
151            return ConvergenceStatus::Unknown;
152        }
153
154        let recent_metrics: Vec<_> =
155            self.metrics_history.iter().rev().take(self.config.convergence_window).collect();
156
157        let losses: Vec<f64> = recent_metrics.iter().map(|m| m.loss).collect();
158
159        // Check for convergence patterns
160        if self.is_converged(&losses) {
161            ConvergenceStatus::Converged
162        } else if self.is_diverging(&losses) {
163            ConvergenceStatus::Diverging
164        } else if self.is_oscillating(&losses) {
165            ConvergenceStatus::Oscillating
166        } else if self.is_plateau(&losses) {
167            ConvergenceStatus::Plateau
168        } else if self.is_converging(&losses) {
169            ConvergenceStatus::Converging
170        } else {
171            ConvergenceStatus::Unknown
172        }
173    }
174
175    /// Assess training stability.
176    pub fn assess_training_stability(&self) -> TrainingStability {
177        if self.metrics_history.len() < 10 {
178            return TrainingStability::Unknown;
179        }
180
181        let recent_losses: Vec<f64> =
182            self.metrics_history.iter().rev().take(20).map(|m| m.loss).collect();
183
184        let variance = self.calculate_variance(&recent_losses);
185
186        if variance > self.config.max_variance_threshold {
187            TrainingStability::Unstable
188        } else if variance > self.config.max_variance_threshold / 2.0 {
189            TrainingStability::HighVariance
190        } else {
191            TrainingStability::Stable
192        }
193    }
194
195    /// Calculate learning efficiency score.
196    pub fn calculate_learning_efficiency(&self) -> f64 {
197        if self.metrics_history.len() < 2 {
198            return 0.0;
199        }
200
201        let initial_loss = self.metrics_history.front().map(|m| m.loss).unwrap_or(0.0);
202        let current_loss = self.metrics_history.back().map(|m| m.loss).unwrap_or(0.0);
203        let steps = self.metrics_history.len();
204
205        if initial_loss <= current_loss {
206            return 0.0;
207        }
208
209        let improvement = (initial_loss - current_loss) / initial_loss;
210        let efficiency = improvement / (steps as f64).sqrt();
211
212        efficiency.min(1.0)
213    }
214
215    /// Detect overfitting indicators from the recorded training metrics.
216    ///
217    /// [`ModelPerformanceMetrics`] carries no validation split, so the
218    /// validation-flavoured variants of [`OverfittingIndicator`]
219    /// (`TrainValidationGap`, `ValidationLossIncreasing`,
220    /// `HighVarianceInValidation`) are never raised here -- they exist for
221    /// callers that really do hold validation data. The two signals this can
222    /// honestly report are a collapsed *training* loss and high variance of the
223    /// *training* loss.
224    ///
225    /// This previously pushed `PerfectTrainingAccuracy` whenever the mean
226    /// training LOSS fell below `0.01` (loss is not accuracy) and
227    /// `HighVarianceInValidation` from the variance of the TRAINING loss (there
228    /// is no validation series to take a variance of).
229    pub fn detect_overfitting_indicators(&self) -> Vec<OverfittingIndicator> {
230        let mut indicators = Vec::new();
231
232        if self.metrics_history.len() > 10 {
233            let recent: Vec<&ModelPerformanceMetrics> =
234                self.metrics_history.iter().rev().take(10).collect();
235            let recent_losses: Vec<f64> = recent.iter().map(|m| m.loss).collect();
236
237            let avg_loss = recent_losses.iter().sum::<f64>() / recent_losses.len() as f64;
238            if avg_loss < 0.01 {
239                indicators.push(OverfittingIndicator::NearZeroTrainingLoss { loss: avg_loss });
240            }
241
242            // Real accuracy, when the caller actually recorded one.
243            let accuracies: Vec<f64> = recent.iter().filter_map(|m| m.accuracy).collect();
244            if !accuracies.is_empty() {
245                let avg_accuracy = accuracies.iter().sum::<f64>() / accuracies.len() as f64;
246                if avg_accuracy >= 0.999 {
247                    indicators.push(OverfittingIndicator::PerfectTrainingAccuracy {
248                        accuracy: avg_accuracy,
249                    });
250                }
251            }
252
253            let variance = self.calculate_variance(&recent_losses);
254            if variance > 0.05 {
255                indicators.push(OverfittingIndicator::HighVarianceInTrainingLoss { variance });
256            }
257        }
258
259        indicators
260    }
261
262    /// Detect underfitting indicators.
263    pub fn detect_underfitting_indicators(&self) -> Vec<UnderfittingIndicator> {
264        let mut indicators = Vec::new();
265
266        if let Some(current_metrics) = self.metrics_history.back() {
267            // High training loss
268            if current_metrics.loss > 1.0 {
269                indicators.push(UnderfittingIndicator::HighTrainingLoss {
270                    loss: current_metrics.loss,
271                    threshold: 1.0,
272                });
273            }
274
275            // Real recorded accuracy, when the caller supplied one.
276            if let Some(accuracy) = current_metrics.accuracy {
277                if accuracy < 0.5 {
278                    indicators.push(UnderfittingIndicator::LowTrainingAccuracy {
279                        accuracy,
280                        threshold: 0.5,
281                    });
282                }
283            }
284
285            // Slow convergence
286            if self.current_state.steps_since_improvement > 50 {
287                indicators.push(UnderfittingIndicator::SlowConvergence {
288                    steps_taken: self.metrics_history.len(),
289                    expected: self.metrics_history.len() / 2,
290                });
291            }
292
293            // No learning
294            if self.current_state.steps_since_improvement > 100 {
295                indicators.push(UnderfittingIndicator::NoLearning {
296                    steps_without_improvement: self.current_state.steps_since_improvement,
297                });
298            }
299        }
300
301        indicators
302    }
303
304    /// Detect plateau in training.
305    pub fn detect_plateau(&self) -> Option<PlateauInfo> {
306        if self.metrics_history.len() < self.config.min_plateau_duration {
307            return None;
308        }
309
310        let recent_losses: Vec<f64> = self
311            .metrics_history
312            .iter()
313            .rev()
314            .take(self.config.min_plateau_duration)
315            .map(|m| m.loss)
316            .collect();
317
318        let variance = self.calculate_variance(&recent_losses);
319        let mean_loss = recent_losses.iter().sum::<f64>() / recent_losses.len() as f64;
320
321        // Check if variance is low enough to indicate plateau
322        if variance < self.config.min_improvement_threshold {
323            let start_step = self.metrics_history.len() - self.config.min_plateau_duration;
324            Some(PlateauInfo {
325                start_step,
326                duration_steps: self.config.min_plateau_duration,
327                plateau_value: mean_loss,
328                variance,
329            })
330        } else {
331            None
332        }
333    }
334
335    /// Generate training recommendations based on current dynamics.
336    pub fn generate_training_recommendations(&self) -> Vec<TrainingRecommendation> {
337        let mut recommendations = Vec::new();
338        let dynamics = self.analyze_training_dynamics();
339
340        match dynamics.convergence_status {
341            ConvergenceStatus::Diverging => {
342                recommendations.push(TrainingRecommendation {
343                    category: "Convergence".to_string(),
344                    priority: TrainingRecommendationPriority::Critical,
345                    description: "Training is diverging".to_string(),
346                    action: "Reduce learning rate immediately".to_string(),
347                    expected_impact: 0.8,
348                });
349            },
350            ConvergenceStatus::Plateau => {
351                recommendations.push(TrainingRecommendation {
352                    category: "Convergence".to_string(),
353                    priority: TrainingRecommendationPriority::High,
354                    description: "Training has reached a plateau".to_string(),
355                    action: "Consider learning rate scheduling or data augmentation".to_string(),
356                    expected_impact: 0.6,
357                });
358            },
359            _ => {},
360        }
361
362        if let TrainingStability::Unstable = dynamics.training_stability {
363            recommendations.push(TrainingRecommendation {
364                category: "Stability".to_string(),
365                priority: TrainingRecommendationPriority::High,
366                description: "Training is unstable".to_string(),
367                action: "Reduce learning rate or add gradient clipping".to_string(),
368                expected_impact: 0.7,
369            });
370        }
371
372        if dynamics.learning_efficiency < 0.3 {
373            recommendations.push(TrainingRecommendation {
374                category: "Efficiency".to_string(),
375                priority: TrainingRecommendationPriority::Medium,
376                description: "Low learning efficiency detected".to_string(),
377                action: "Consider architecture changes or hyperparameter tuning".to_string(),
378                expected_impact: 0.5,
379            });
380        }
381
382        recommendations
383    }
384
385    // Helper methods for convergence detection
386    fn is_converged(&self, losses: &[f64]) -> bool {
387        if losses.len() < 5 {
388            return false;
389        }
390
391        let recent_variance = self.calculate_variance(&losses[..5]);
392        recent_variance < self.config.min_improvement_threshold && losses[0] < 0.01
393    }
394
395    fn is_diverging(&self, losses: &[f64]) -> bool {
396        if losses.len() < 3 {
397            return false;
398        }
399
400        let (Some(&first), Some(&last)) = (losses.first(), losses.last()) else {
401            return false;
402        };
403        // Check if loss is consistently increasing
404        losses.windows(2).all(|w| w[1] >= w[0]) && (last / first) > 1.1
405    }
406
407    fn is_oscillating(&self, losses: &[f64]) -> bool {
408        if losses.len() < 6 {
409            return false;
410        }
411
412        // Check for oscillating pattern
413        let mut direction_changes = 0;
414        for window in losses.windows(3) {
415            let trend1 = window[1] - window[0];
416            let trend2 = window[2] - window[1];
417            if trend1.signum() != trend2.signum() {
418                direction_changes += 1;
419            }
420        }
421
422        direction_changes > losses.len() / 3
423    }
424
425    fn is_plateau(&self, losses: &[f64]) -> bool {
426        let variance = self.calculate_variance(losses);
427        variance < self.config.min_improvement_threshold
428    }
429
430    fn is_converging(&self, losses: &[f64]) -> bool {
431        if losses.len() < 3 {
432            return false;
433        }
434
435        // Check if loss is generally decreasing
436        let trend = self.calculate_trend(losses);
437        trend < -self.config.min_improvement_threshold
438    }
439
440    fn calculate_variance(&self, values: &[f64]) -> f64 {
441        if values.len() < 2 {
442            return 0.0;
443        }
444
445        let mean = values.iter().sum::<f64>() / values.len() as f64;
446        let variance =
447            values.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / (values.len() - 1) as f64;
448        variance
449    }
450
451    fn calculate_trend(&self, values: &[f64]) -> f64 {
452        if values.len() < 2 {
453            return 0.0;
454        }
455
456        let n = values.len() as f64;
457        let x_mean = (n - 1.0) / 2.0;
458        let y_mean = values.iter().sum::<f64>() / n;
459
460        let mut numerator = 0.0;
461        let mut denominator = 0.0;
462
463        for (i, &y) in values.iter().enumerate() {
464            let x = i as f64;
465            numerator += (x - x_mean) * (y - y_mean);
466            denominator += (x - x_mean).powi(2);
467        }
468
469        if denominator == 0.0 {
470            0.0
471        } else {
472            numerator / denominator
473        }
474    }
475
476    /// Clear analysis history.
477    pub fn clear(&mut self) {
478        self.metrics_history.clear();
479        self.current_state = TrainingState::default();
480    }
481
482    /// Get current training state information.
483    pub fn get_training_state(&self) -> &TrainingState {
484        &self.current_state
485    }
486
487    /// Generate comprehensive training dynamics report.
488    pub async fn generate_report(&self) -> anyhow::Result<TrainingDynamicsReport> {
489        let training_dynamics = self.analyze_training_dynamics();
490        let recommendations = self.generate_recommendations();
491
492        Ok(TrainingDynamicsReport {
493            training_dynamics,
494            recommendations,
495            current_state: self.current_state.clone(),
496        })
497    }
498
499    /// Generate training recommendations.
500    fn generate_recommendations(&self) -> Vec<TrainingRecommendation> {
501        let mut recommendations = Vec::new();
502
503        // Add basic recommendations based on current state
504        recommendations.push(TrainingRecommendation {
505            category: "General".to_string(),
506            description: "Continue monitoring training dynamics".to_string(),
507            action: "Monitor training progress and adjust parameters as needed".to_string(),
508            priority: TrainingRecommendationPriority::Low,
509            expected_impact: 0.1,
510        });
511
512        recommendations
513    }
514}
515
516impl Default for TrainingDynamicsAnalyzer {
517    fn default() -> Self {
518        Self::new()
519    }
520}
521
522/// Training recommendation.
523#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
524pub struct TrainingRecommendation {
525    /// Category of the recommendation
526    pub category: String,
527    /// Priority level
528    pub priority: TrainingRecommendationPriority,
529    /// Description of the issue
530    pub description: String,
531    /// Recommended action
532    pub action: String,
533    /// Expected impact (0.0 to 1.0)
534    pub expected_impact: f64,
535}
536
537/// Priority levels for training recommendations.
538#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
539pub enum TrainingRecommendationPriority {
540    /// Low priority
541    Low,
542    /// Medium priority
543    Medium,
544    /// High priority
545    High,
546    /// Critical priority
547    Critical,
548}
549
550/// Comprehensive training dynamics report.
551#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
552pub struct TrainingDynamicsReport {
553    /// Training dynamics analysis
554    pub training_dynamics: TrainingDynamics,
555    /// Generated recommendations
556    pub recommendations: Vec<TrainingRecommendation>,
557    /// Current training state
558    pub current_state: TrainingState,
559}
560
561#[cfg(test)]
562mod tests {
563    use super::*;
564    use chrono::Utc;
565
566    fn create_test_metrics(step: usize, loss: f64) -> ModelPerformanceMetrics {
567        ModelPerformanceMetrics {
568            training_step: step,
569            loss,
570            accuracy: Some(0.8),
571            learning_rate: 0.001,
572            batch_size: 32,
573            throughput_samples_per_sec: 100.0,
574            memory_usage_mb: 1000.0,
575            gpu_utilization: Some(0.9),
576            timestamp: Utc::now(),
577        }
578    }
579
580    #[test]
581    fn test_training_dynamics_analyzer_creation() {
582        let analyzer = TrainingDynamicsAnalyzer::new();
583        assert_eq!(analyzer.metrics_history.len(), 0);
584    }
585
586    #[test]
587    fn test_add_metrics() {
588        let mut analyzer = TrainingDynamicsAnalyzer::new();
589        let metrics = create_test_metrics(1, 0.5);
590
591        analyzer.add_metrics(metrics);
592        assert_eq!(analyzer.metrics_history.len(), 1);
593        assert_eq!(analyzer.current_state.best_loss, 0.5);
594    }
595
596    #[test]
597    fn test_convergence_detection() {
598        let mut analyzer = TrainingDynamicsAnalyzer::new();
599
600        // Add converging sequence
601        for i in 1..=25 {
602            let loss = 1.0 / (i as f64);
603            let metrics = create_test_metrics(i, loss);
604            analyzer.add_metrics(metrics);
605        }
606
607        let status = analyzer.detect_convergence_status();
608        matches!(
609            status,
610            ConvergenceStatus::Converging | ConvergenceStatus::Converged
611        );
612    }
613
614    #[test]
615    fn test_learning_efficiency_calculation() {
616        let mut analyzer = TrainingDynamicsAnalyzer::new();
617
618        analyzer.add_metrics(create_test_metrics(1, 1.0));
619        analyzer.add_metrics(create_test_metrics(2, 0.5));
620        analyzer.add_metrics(create_test_metrics(3, 0.25));
621
622        let efficiency = analyzer.calculate_learning_efficiency();
623        assert!(efficiency > 0.0);
624    }
625
626    #[test]
627    fn test_plateau_detection() {
628        let mut analyzer = TrainingDynamicsAnalyzer::new();
629
630        // Add plateau sequence
631        for i in 1..=15 {
632            let metrics = create_test_metrics(i, 0.1); // Constant loss
633            analyzer.add_metrics(metrics);
634        }
635
636        let plateau = analyzer.detect_plateau();
637        assert!(plateau.is_some());
638    }
639
640    #[test]
641    fn test_training_stability_assessment() {
642        let mut analyzer = TrainingDynamicsAnalyzer::new();
643
644        // Add stable sequence
645        for i in 1..=20 {
646            let loss = 0.5 + (i as f64 * 0.001); // Very small variance
647            let metrics = create_test_metrics(i, loss);
648            analyzer.add_metrics(metrics);
649        }
650
651        let stability = analyzer.assess_training_stability();
652        matches!(stability, TrainingStability::Stable);
653    }
654}