use crate::{OptimizerError, OptimizerResult, OptimizerState};
use parking_lot::RwLock;
use std::collections::{HashMap, VecDeque};
use std::sync::Arc;
use torsh_core::error::Result;
use torsh_tensor::Tensor;
pub struct OptimizerAnalyzer {
step_history: Vec<OptimizationStep>,
gradient_stats: GradientStatistics,
parameter_stats: ParameterStatistics,
convergence_tracker: ConvergenceTracker,
config: AnalyzerConfig,
}
#[derive(Debug, Clone)]
pub struct AnalyzerConfig {
pub max_history_size: usize,
pub track_gradient_norms: bool,
pub track_parameter_norms: bool,
pub track_gradient_flow: bool,
pub moving_average_window: usize,
}
impl Default for AnalyzerConfig {
fn default() -> Self {
Self {
max_history_size: 10000,
track_gradient_norms: true,
track_parameter_norms: true,
track_gradient_flow: true,
moving_average_window: 100,
}
}
}
#[derive(Debug, Clone)]
pub struct OptimizationStep {
pub step: usize,
pub learning_rates: Vec<f32>,
pub gradient_norms: Vec<f32>,
pub parameter_norms: Vec<f32>,
pub update_norms: Vec<f32>,
pub loss: Option<f32>,
pub timestamp: std::time::Instant,
}
#[derive(Debug, Clone)]
pub struct GradientStatistics {
pub norm_history: VecDeque<f32>,
pub average_norm: f32,
pub max_norm: f32,
pub min_norm: f32,
pub explosion_count: usize,
pub vanishing_count: usize,
pub recent_variance: f32,
}
#[derive(Debug, Clone)]
pub struct ParameterStatistics {
pub norm_history: VecDeque<f32>,
pub update_ratios: VecDeque<f32>,
pub average_update_ratio: f32,
pub velocity: f32,
pub stability_score: f32,
}
#[derive(Debug, Clone)]
pub struct ConvergenceTracker {
pub loss_history: VecDeque<f32>,
pub loss_moving_average: f32,
pub convergence_rate: f32,
pub is_converged: bool,
pub steps_since_improvement: usize,
pub best_loss: f32,
pub plateau_length: usize,
}
impl OptimizerAnalyzer {
pub fn new(config: Option<AnalyzerConfig>) -> Self {
let config = config.unwrap_or_default();
let max_size = config.moving_average_window;
Self {
step_history: Vec::new(),
gradient_stats: GradientStatistics {
norm_history: VecDeque::with_capacity(max_size),
average_norm: 0.0,
max_norm: 0.0,
min_norm: f32::INFINITY,
explosion_count: 0,
vanishing_count: 0,
recent_variance: 0.0,
},
parameter_stats: ParameterStatistics {
norm_history: VecDeque::with_capacity(max_size),
update_ratios: VecDeque::with_capacity(max_size),
average_update_ratio: 0.0,
velocity: 0.0,
stability_score: 1.0,
},
convergence_tracker: ConvergenceTracker {
loss_history: VecDeque::with_capacity(max_size),
loss_moving_average: 0.0,
convergence_rate: 0.0,
is_converged: false,
steps_since_improvement: 0,
best_loss: f32::INFINITY,
plateau_length: 0,
},
config,
}
}
pub fn analyze_step(
&mut self,
step: usize,
params: &[Arc<RwLock<Tensor>>],
learning_rates: &[f32],
loss: Option<f32>,
) -> Result<()> {
let timestamp = std::time::Instant::now();
let mut gradient_norms = Vec::new();
let mut parameter_norms = Vec::new();
let mut update_norms = Vec::new();
for param in params {
let param_tensor = param.read();
let param_norm = param_tensor.norm()?;
parameter_norms.push(param_norm.item()?);
if let Some(grad) = param_tensor.grad() {
let grad_norm = grad.norm()?;
gradient_norms.push(grad_norm.item()?);
let lr = learning_rates.get(0).copied().unwrap_or(0.001);
update_norms.push(lr * grad_norm.item()?);
} else {
gradient_norms.push(0.0);
update_norms.push(0.0);
}
}
let opt_step = OptimizationStep {
step,
learning_rates: learning_rates.to_vec(),
gradient_norms: gradient_norms.clone(),
parameter_norms: parameter_norms.clone(),
update_norms: update_norms.clone(),
loss,
timestamp,
};
self.step_history.push(opt_step);
if self.step_history.len() > self.config.max_history_size {
self.step_history.remove(0);
}
if self.config.track_gradient_norms {
self.update_gradient_stats(&gradient_norms)?;
}
if self.config.track_parameter_norms {
self.update_parameter_stats(¶meter_norms, &update_norms)?;
}
if let Some(loss_val) = loss {
self.update_convergence_tracking(loss_val)?;
}
Ok(())
}
fn update_gradient_stats(&mut self, gradient_norms: &[f32]) -> Result<()> {
for &norm in gradient_norms {
self.gradient_stats.norm_history.push_back(norm);
if self.gradient_stats.norm_history.len() > self.config.moving_average_window {
self.gradient_stats.norm_history.pop_front();
}
self.gradient_stats.max_norm = self.gradient_stats.max_norm.max(norm);
self.gradient_stats.min_norm = self.gradient_stats.min_norm.min(norm);
if norm > 10.0 {
self.gradient_stats.explosion_count += 1;
}
if norm < 1e-7 {
self.gradient_stats.vanishing_count += 1;
}
}
if !self.gradient_stats.norm_history.is_empty() {
let sum: f32 = self.gradient_stats.norm_history.iter().sum();
self.gradient_stats.average_norm = sum / self.gradient_stats.norm_history.len() as f32;
let mean = self.gradient_stats.average_norm;
let variance: f32 = self
.gradient_stats
.norm_history
.iter()
.map(|&x| (x - mean).powi(2))
.sum::<f32>()
/ self.gradient_stats.norm_history.len() as f32;
self.gradient_stats.recent_variance = variance;
}
Ok(())
}
fn update_parameter_stats(
&mut self,
parameter_norms: &[f32],
update_norms: &[f32],
) -> Result<()> {
for (param_norm, update_norm) in parameter_norms.iter().zip(update_norms.iter()) {
self.parameter_stats.norm_history.push_back(*param_norm);
if self.parameter_stats.norm_history.len() > self.config.moving_average_window {
self.parameter_stats.norm_history.pop_front();
}
let ratio = if *param_norm > 1e-8 {
update_norm / param_norm
} else {
0.0
};
self.parameter_stats.update_ratios.push_back(ratio);
if self.parameter_stats.update_ratios.len() > self.config.moving_average_window {
self.parameter_stats.update_ratios.pop_front();
}
}
if !self.parameter_stats.update_ratios.is_empty() {
let sum: f32 = self.parameter_stats.update_ratios.iter().sum();
self.parameter_stats.average_update_ratio =
sum / self.parameter_stats.update_ratios.len() as f32;
}
if self.parameter_stats.norm_history.len() >= 2 {
let current = self
.parameter_stats
.norm_history
.back()
.expect("norm_history should not be empty");
let previous = self
.parameter_stats
.norm_history
.get(self.parameter_stats.norm_history.len() - 2)
.expect("second-to-last element should exist");
self.parameter_stats.velocity = (current - previous).abs();
}
if self.parameter_stats.update_ratios.len() > 1 {
let mean = self.parameter_stats.average_update_ratio;
let variance: f32 = self
.parameter_stats
.update_ratios
.iter()
.map(|&x| (x - mean).powi(2))
.sum::<f32>()
/ self.parameter_stats.update_ratios.len() as f32;
self.parameter_stats.stability_score = 1.0 / (1.0 + variance);
}
Ok(())
}
fn update_convergence_tracking(&mut self, loss: f32) -> Result<()> {
self.convergence_tracker.loss_history.push_back(loss);
if self.convergence_tracker.loss_history.len() > self.config.moving_average_window {
self.convergence_tracker.loss_history.pop_front();
}
if !self.convergence_tracker.loss_history.is_empty() {
let sum: f32 = self.convergence_tracker.loss_history.iter().sum();
self.convergence_tracker.loss_moving_average =
sum / self.convergence_tracker.loss_history.len() as f32;
}
if loss < self.convergence_tracker.best_loss {
self.convergence_tracker.best_loss = loss;
self.convergence_tracker.steps_since_improvement = 0;
self.convergence_tracker.plateau_length = 0;
} else {
self.convergence_tracker.steps_since_improvement += 1;
self.convergence_tracker.plateau_length += 1;
}
if self.convergence_tracker.loss_history.len() >= 10 {
let recent_losses: Vec<f32> = self
.convergence_tracker
.loss_history
.iter()
.rev()
.take(10)
.cloned()
.collect();
let n = recent_losses.len() as f32;
let x_mean = (n - 1.0) / 2.0;
let y_mean = recent_losses.iter().sum::<f32>() / n;
let numerator: f32 = recent_losses
.iter()
.enumerate()
.map(|(i, &y)| (i as f32 - x_mean) * (y - y_mean))
.sum();
let denominator: f32 = (0..recent_losses.len())
.map(|i| (i as f32 - x_mean).powi(2))
.sum();
if denominator > 1e-8 {
self.convergence_tracker.convergence_rate = numerator / denominator;
}
}
self.convergence_tracker.is_converged = self.convergence_tracker.plateau_length > 1000
|| (self.convergence_tracker.convergence_rate.abs() < 1e-6
&& self.convergence_tracker.loss_history.len() > 100);
Ok(())
}
pub fn generate_report(&self) -> AnalysisReport {
AnalysisReport {
total_steps: self.step_history.len(),
gradient_stats: self.gradient_stats.clone(),
parameter_stats: self.parameter_stats.clone(),
convergence_tracker: self.convergence_tracker.clone(),
recommendations: self.generate_recommendations(),
}
}
fn generate_recommendations(&self) -> Vec<OptimizationRecommendation> {
let mut recommendations = Vec::new();
if self.gradient_stats.explosion_count > 10 {
recommendations.push(OptimizationRecommendation {
category: RecommendationCategory::GradientNorms,
severity: Severity::High,
message: "Frequent gradient explosions detected. Consider gradient clipping or reducing learning rate.".to_string(),
suggested_actions: vec![
"Add gradient clipping with max_norm=1.0".to_string(),
"Reduce learning rate by factor of 10".to_string(),
"Use adaptive optimizers like Adam".to_string(),
],
});
}
if self.gradient_stats.vanishing_count > 10 {
recommendations.push(OptimizationRecommendation {
category: RecommendationCategory::GradientNorms,
severity: Severity::Medium,
message: "Frequent gradient vanishing detected. Model may be too deep or have saturation issues.".to_string(),
suggested_actions: vec![
"Check activation functions for saturation".to_string(),
"Consider batch normalization".to_string(),
"Use residual connections".to_string(),
],
});
}
if self.parameter_stats.average_update_ratio > 0.1 {
recommendations.push(OptimizationRecommendation {
category: RecommendationCategory::LearningRate,
severity: Severity::Medium,
message: "Update ratios are high. Learning rate might be too large.".to_string(),
suggested_actions: vec![
"Reduce learning rate by factor of 2-5".to_string(),
"Use learning rate scheduling".to_string(),
],
});
} else if self.parameter_stats.average_update_ratio < 0.001 {
recommendations.push(OptimizationRecommendation {
category: RecommendationCategory::LearningRate,
severity: Severity::Low,
message: "Update ratios are very small. Learning rate might be too small."
.to_string(),
suggested_actions: vec![
"Increase learning rate by factor of 2-10".to_string(),
"Consider warmup schedule".to_string(),
],
});
}
if self.convergence_tracker.plateau_length > 500 {
recommendations.push(OptimizationRecommendation {
category: RecommendationCategory::Convergence,
severity: Severity::Medium,
message: "Training has plateaued for many steps.".to_string(),
suggested_actions: vec![
"Reduce learning rate".to_string(),
"Add regularization".to_string(),
"Consider early stopping".to_string(),
],
});
}
recommendations
}
pub fn get_gradient_flow_data(&self, num_steps: usize) -> Vec<GradientFlowPoint> {
self.step_history
.iter()
.rev()
.take(num_steps)
.map(|step| GradientFlowPoint {
step: step.step,
gradient_norms: step.gradient_norms.clone(),
parameter_norms: step.parameter_norms.clone(),
update_norms: step.update_norms.clone(),
loss: step.loss,
})
.collect()
}
}
#[derive(Debug, Clone)]
pub struct AnalysisReport {
pub total_steps: usize,
pub gradient_stats: GradientStatistics,
pub parameter_stats: ParameterStatistics,
pub convergence_tracker: ConvergenceTracker,
pub recommendations: Vec<OptimizationRecommendation>,
}
#[derive(Debug, Clone)]
pub struct OptimizationRecommendation {
pub category: RecommendationCategory,
pub severity: Severity,
pub message: String,
pub suggested_actions: Vec<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum RecommendationCategory {
GradientNorms,
LearningRate,
Convergence,
Stability,
Performance,
}
#[derive(Debug, Clone, PartialEq)]
pub enum Severity {
Low,
Medium,
High,
Critical,
}
#[derive(Debug, Clone)]
pub struct GradientFlowPoint {
pub step: usize,
pub gradient_norms: Vec<f32>,
pub parameter_norms: Vec<f32>,
pub update_norms: Vec<f32>,
pub loss: Option<f32>,
}
pub struct HyperparameterSensitivity {
sensitivity_data: HashMap<String, SensitivityResult>,
base_config: HashMap<String, f32>,
}
#[derive(Debug, Clone)]
pub struct SensitivityResult {
pub name: String,
pub test_values: Vec<f32>,
pub performance_metrics: Vec<f32>,
pub sensitivity_score: f32,
pub optimal_value: f32,
}
impl HyperparameterSensitivity {
pub fn new(base_config: HashMap<String, f32>) -> Self {
Self {
sensitivity_data: HashMap::new(),
base_config,
}
}
pub fn analyze_parameter(
&mut self,
param_name: &str,
test_values: Vec<f32>,
performance_evaluator: impl Fn(f32) -> Result<f32>,
) -> Result<SensitivityResult> {
let mut performance_metrics = Vec::new();
for &value in &test_values {
let performance = performance_evaluator(value)?;
performance_metrics.push(performance);
}
let max_perf = performance_metrics
.iter()
.fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let min_perf = performance_metrics
.iter()
.fold(f32::INFINITY, |a, &b| a.min(b));
let sensitivity_score = if max_perf != min_perf {
(max_perf - min_perf) / max_perf.abs()
} else {
0.0
};
let optimal_idx = performance_metrics
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(i, _)| i)
.unwrap_or(0);
let optimal_value = test_values[optimal_idx];
let result = SensitivityResult {
name: param_name.to_string(),
test_values,
performance_metrics,
sensitivity_score,
optimal_value,
};
self.sensitivity_data
.insert(param_name.to_string(), result.clone());
Ok(result)
}
pub fn get_sensitivity_ranking(&self) -> Vec<(&str, f32)> {
let mut ranking: Vec<_> = self
.sensitivity_data
.iter()
.map(|(name, result)| (name.as_str(), result.sensitivity_score))
.collect();
ranking.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
ranking
}
pub fn generate_sensitivity_report(&self) -> SensitivityReport {
let ranking = self.get_sensitivity_ranking();
let most_sensitive = ranking
.first()
.map(|(name, score)| (name.to_string(), *score));
let least_sensitive = ranking
.last()
.map(|(name, score)| (name.to_string(), *score));
SensitivityReport {
analyzed_parameters: self.sensitivity_data.keys().cloned().collect(),
sensitivity_ranking: ranking
.into_iter()
.map(|(n, s)| (n.to_string(), s))
.collect(),
most_sensitive_parameter: most_sensitive,
least_sensitive_parameter: least_sensitive,
recommendations: self.generate_sensitivity_recommendations(),
}
}
fn generate_sensitivity_recommendations(&self) -> Vec<String> {
let mut recommendations = Vec::new();
let ranking = self.get_sensitivity_ranking();
if let Some((most_sensitive, score)) = ranking.first() {
if *score > 0.5 {
recommendations.push(format!(
"Parameter '{}' is highly sensitive (score: {:.3}). Fine-tune carefully.",
most_sensitive, score
));
}
}
if let Some((least_sensitive, score)) = ranking.last() {
if *score < 0.1 {
recommendations.push(format!(
"Parameter '{}' has low sensitivity (score: {:.3}). Consider using default values.",
least_sensitive, score
));
}
}
recommendations
}
}
#[derive(Debug, Clone)]
pub struct SensitivityReport {
pub analyzed_parameters: Vec<String>,
pub sensitivity_ranking: Vec<(String, f32)>,
pub most_sensitive_parameter: Option<(String, f32)>,
pub least_sensitive_parameter: Option<(String, f32)>,
pub recommendations: Vec<String>,
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_tensor::creation::randn;
#[test]
fn test_optimizer_analyzer_creation() {
let analyzer = OptimizerAnalyzer::new(None);
assert_eq!(analyzer.step_history.len(), 0);
}
#[test]
fn test_analyzer_step() {
let mut analyzer = OptimizerAnalyzer::new(None);
let params = vec![Arc::new(RwLock::new(randn::<f32>(&[10, 10]).unwrap()))];
let result = analyzer.analyze_step(1, ¶ms, &[0.01], Some(0.5));
assert!(result.is_ok());
assert_eq!(analyzer.step_history.len(), 1);
}
#[test]
fn test_sensitivity_analyzer() -> OptimizerResult<()> {
let base_config = [("lr".to_string(), 0.01)].iter().cloned().collect();
let mut sensitivity = HyperparameterSensitivity::new(base_config);
let test_values = vec![0.001, 0.01, 0.1];
let evaluator = |lr: f32| Ok(1.0 / lr);
let _result = sensitivity.analyze_parameter("lr", test_values, evaluator)?;
let ranking = sensitivity.get_sensitivity_ranking();
assert_eq!(ranking.len(), 1);
assert_eq!(ranking[0].0, "lr");
Ok(())
}
}