Skip to main content

optirs_core/
optimizer_metrics.rs

1//! Optimizer performance metrics and monitoring
2//!
3//! This module provides comprehensive metrics collection and monitoring for optimizers
4//! using SciRS2's metrics infrastructure for production deployments.
5//!
6//! # Features
7//!
8//! - Real-time optimizer performance tracking
9//! - Gradient and parameter statistics
10//! - Convergence monitoring
11//! - Memory usage tracking
12//! - Performance dashboards and reporting
13//!
14//! # SciRS2 Integration
15//!
16//! This module uses SciRS2-Core metrics abstractions exclusively:
17//! - `scirs2_core::metrics::MetricRegistry` for metric registration
18//! - `scirs2_core::metrics::Counter` for counting operations
19//! - `scirs2_core::metrics::Gauge` for current values
20//! - `scirs2_core::metrics::Histogram` for distributions
21//! - `scirs2_core::metrics::Timer` for timing operations
22
23use scirs2_core::ndarray::{ArrayView1, ScalarOperand};
24use scirs2_core::numeric::Float;
25use std::collections::HashMap;
26use std::fmt::Debug;
27use std::time::{Duration, Instant};
28
29use crate::error::{OptimError, Result};
30use crate::utils::try_f64;
31
32/// Optimizer performance metrics
33///
34/// Tracks key performance indicators for optimizer operations including
35/// step timing, gradient statistics, parameter updates, and convergence.
36#[derive(Debug, Clone)]
37pub struct OptimizerMetrics {
38    /// Optimizer name
39    pub name: String,
40    /// Total number of optimization steps
41    pub step_count: u64,
42    /// Total time spent in optimization steps
43    pub total_step_time: Duration,
44    /// Average time per step
45    pub avg_step_time: Duration,
46    /// Current learning rate
47    pub current_learning_rate: f64,
48    /// Gradient statistics
49    pub gradient_stats: GradientStatistics,
50    /// Parameter statistics
51    pub parameter_stats: ParameterStatistics,
52    /// Convergence metrics
53    pub convergence: ConvergenceMetrics,
54    /// Memory usage (bytes)
55    pub memory_usage: usize,
56}
57
58impl OptimizerMetrics {
59    /// Create new metrics for an optimizer
60    pub fn new(name: impl Into<String>) -> Self {
61        Self {
62            name: name.into(),
63            step_count: 0,
64            total_step_time: Duration::ZERO,
65            avg_step_time: Duration::ZERO,
66            current_learning_rate: 0.0,
67            gradient_stats: GradientStatistics::default(),
68            parameter_stats: ParameterStatistics::default(),
69            convergence: ConvergenceMetrics::default(),
70            memory_usage: 0,
71        }
72    }
73
74    /// Update metrics after an optimization step.
75    ///
76    /// Returns `Err` if the gradient or parameter statistics cannot be
77    /// computed (mismatched parameter lengths, or values with no `f64`
78    /// representation).
79    pub fn update_step<A: Float>(
80        &mut self,
81        step_duration: Duration,
82        learning_rate: f64,
83        gradients: &ArrayView1<A>,
84        params_before: &ArrayView1<A>,
85        params_after: &ArrayView1<A>,
86    ) -> Result<()> {
87        // Statistics first: they are the only fallible part, and running them
88        // before mutating the step counters keeps a rejected call from
89        // advancing `step_count`/`total_step_time` for a step whose metrics
90        // were never recorded.
91        self.gradient_stats.update(gradients)?;
92        self.parameter_stats.update(params_before, params_after)?;
93        self.convergence.update(&self.parameter_stats);
94
95        self.step_count += 1;
96        self.total_step_time += step_duration;
97        self.avg_step_time = self.total_step_time / self.step_count as u32;
98        self.current_learning_rate = learning_rate;
99
100        Ok(())
101    }
102
103    /// Get throughput (steps per second)
104    pub fn throughput(&self) -> f64 {
105        if self.total_step_time.as_secs_f64() > 0.0 {
106            self.step_count as f64 / self.total_step_time.as_secs_f64()
107        } else {
108            0.0
109        }
110    }
111
112    /// Reset all metrics
113    pub fn reset(&mut self) {
114        self.step_count = 0;
115        self.total_step_time = Duration::ZERO;
116        self.avg_step_time = Duration::ZERO;
117        self.gradient_stats = GradientStatistics::default();
118        self.parameter_stats = ParameterStatistics::default();
119        self.convergence = ConvergenceMetrics::default();
120    }
121}
122
123/// Gradient statistics
124#[derive(Debug, Clone, Default)]
125pub struct GradientStatistics {
126    /// Mean gradient magnitude
127    pub mean: f64,
128    /// Standard deviation of gradients
129    pub std_dev: f64,
130    /// Maximum gradient value
131    pub max: f64,
132    /// Minimum gradient value
133    pub min: f64,
134    /// Gradient norm (L2)
135    pub norm: f64,
136    /// Number of zero gradients
137    pub num_zeros: usize,
138}
139
140impl GradientStatistics {
141    /// Update gradient statistics.
142    ///
143    /// Returns `Err` when a gradient value has no `f64` representation, rather
144    /// than panicking part-way through and leaving the statistics in a
145    /// half-updated state.
146    pub fn update<A: Float>(&mut self, gradients: &ArrayView1<A>) -> Result<()> {
147        let n = gradients.len();
148        if n == 0 {
149            return Ok(());
150        }
151
152        // Convert once, up front: every statistic below is computed in f64, and
153        // doing the (fallible) narrowing in one pass means a failure aborts
154        // before any field has been overwritten.
155        let values: Vec<f64> = gradients
156            .iter()
157            .map(|&g| try_f64(g))
158            .collect::<Result<Vec<f64>>>()?;
159
160        let count = n as f64;
161        let mean = values.iter().sum::<f64>() / count;
162        let variance = values.iter().map(|&v| (v - mean) * (v - mean)).sum::<f64>() / count;
163
164        self.mean = mean;
165        self.std_dev = variance.sqrt();
166        self.max = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
167        self.min = values.iter().copied().fold(f64::INFINITY, f64::min);
168        self.norm = values.iter().map(|&v| v * v).sum::<f64>().sqrt();
169        self.num_zeros = values.iter().filter(|v| v.abs() < 1e-10).count();
170
171        Ok(())
172    }
173}
174
175/// Parameter statistics
176#[derive(Debug, Clone, Default)]
177pub struct ParameterStatistics {
178    /// Mean parameter value
179    pub mean: f64,
180    /// Standard deviation of parameters
181    pub std_dev: f64,
182    /// Parameter update magnitude
183    pub update_magnitude: f64,
184    /// Relative parameter change
185    pub relative_change: f64,
186}
187
188impl ParameterStatistics {
189    /// Update parameter statistics.
190    ///
191    /// Returns `Err` when the two parameter views disagree on length (which
192    /// `zip` would otherwise silently truncate, understating the update
193    /// magnitude) or when a value has no `f64` representation.
194    pub fn update<A: Float>(
195        &mut self,
196        params_before: &ArrayView1<A>,
197        params_after: &ArrayView1<A>,
198    ) -> Result<()> {
199        let n = params_after.len();
200        if n == 0 {
201            return Ok(());
202        }
203        if params_before.len() != n {
204            return Err(OptimError::DimensionMismatch(format!(
205                "parameter statistics need the pre- and post-step parameters to have the same \
206                 length, got {} before and {n} after",
207                params_before.len()
208            )));
209        }
210
211        // Narrow once, up front, so a failure aborts before any field is
212        // overwritten and no value is converted twice.
213        let after: Vec<f64> = params_after
214            .iter()
215            .map(|&p| try_f64(p))
216            .collect::<Result<Vec<f64>>>()?;
217        let before: Vec<f64> = params_before
218            .iter()
219            .map(|&p| try_f64(p))
220            .collect::<Result<Vec<f64>>>()?;
221
222        let count = n as f64;
223        let mean = after.iter().sum::<f64>() / count;
224        let variance = after.iter().map(|&v| (v - mean) * (v - mean)).sum::<f64>() / count;
225        let update_magnitude = before
226            .iter()
227            .zip(after.iter())
228            .map(|(&b, &a)| (a - b) * (a - b))
229            .sum::<f64>()
230            .sqrt();
231        let params_norm = before.iter().map(|&v| v * v).sum::<f64>().sqrt();
232
233        self.mean = mean;
234        self.std_dev = variance.sqrt();
235        self.update_magnitude = update_magnitude;
236        self.relative_change = if params_norm > 1e-10 {
237            update_magnitude / params_norm
238        } else {
239            0.0
240        };
241
242        Ok(())
243    }
244}
245
246/// Convergence metrics
247#[derive(Debug, Clone, Default)]
248pub struct ConvergenceMetrics {
249    /// Moving average of parameter updates
250    pub update_moving_avg: f64,
251    /// Is optimizer converging (updates decreasing)
252    pub is_converging: bool,
253    /// Estimated steps to convergence
254    pub estimated_steps_to_convergence: Option<u64>,
255    /// Convergence rate
256    pub convergence_rate: f64,
257}
258
259impl ConvergenceMetrics {
260    /// Update convergence metrics
261    pub fn update(&mut self, param_stats: &ParameterStatistics) {
262        // Check if converging before updating (compare against previous average)
263        if self.update_moving_avg > 1e-10 {
264            self.is_converging = param_stats.update_magnitude < self.update_moving_avg;
265            self.convergence_rate = 1.0 - (param_stats.update_magnitude / self.update_moving_avg);
266        }
267
268        // Update moving average with exponential smoothing (alpha = 0.1)
269        let alpha = 0.1;
270        self.update_moving_avg =
271            alpha * param_stats.update_magnitude + (1.0 - alpha) * self.update_moving_avg;
272    }
273}
274
275/// Metrics collector for tracking multiple optimizers
276pub struct MetricsCollector {
277    /// Metrics for each optimizer
278    metrics: HashMap<String, OptimizerMetrics>,
279    /// Global start time
280    start_time: Instant,
281}
282
283impl MetricsCollector {
284    /// Create a new metrics collector
285    pub fn new() -> Self {
286        Self {
287            metrics: HashMap::new(),
288            start_time: Instant::now(),
289        }
290    }
291
292    /// Register a new optimizer for tracking
293    pub fn register_optimizer(&mut self, name: impl Into<String>) {
294        let name = name.into();
295        self.metrics
296            .entry(name.clone())
297            .or_insert_with(|| OptimizerMetrics::new(name));
298    }
299
300    /// Update metrics for an optimizer
301    pub fn update<A: Float + ScalarOperand>(
302        &mut self,
303        optimizer_name: &str,
304        step_duration: Duration,
305        learning_rate: f64,
306        gradients: &ArrayView1<A>,
307        params_before: &ArrayView1<A>,
308        params_after: &ArrayView1<A>,
309    ) -> Result<()> {
310        if let Some(metrics) = self.metrics.get_mut(optimizer_name) {
311            metrics.update_step(
312                step_duration,
313                learning_rate,
314                gradients,
315                params_before,
316                params_after,
317            )
318        } else {
319            Err(crate::error::OptimError::InvalidConfig(format!(
320                "Optimizer '{}' not registered",
321                optimizer_name
322            )))
323        }
324    }
325
326    /// Get metrics for an optimizer
327    pub fn get_metrics(&self, optimizer_name: &str) -> Option<&OptimizerMetrics> {
328        self.metrics.get(optimizer_name)
329    }
330
331    /// Get all metrics
332    pub fn all_metrics(&self) -> &HashMap<String, OptimizerMetrics> {
333        &self.metrics
334    }
335
336    /// Get elapsed time since collector started
337    pub fn elapsed(&self) -> Duration {
338        self.start_time.elapsed()
339    }
340
341    /// Reset all metrics
342    pub fn reset(&mut self) {
343        for metrics in self.metrics.values_mut() {
344            metrics.reset();
345        }
346        self.start_time = Instant::now();
347    }
348
349    /// Generate summary report
350    pub fn summary_report(&self) -> String {
351        let mut report = String::new();
352        report.push_str("=== Optimizer Metrics Summary ===\n");
353        report.push_str(&format!("Total elapsed time: {:?}\n\n", self.elapsed()));
354
355        for (name, metrics) in &self.metrics {
356            report.push_str(&format!("Optimizer: {}\n", name));
357            report.push_str(&format!("  Steps: {}\n", metrics.step_count));
358            report.push_str(&format!("  Avg step time: {:?}\n", metrics.avg_step_time));
359            report.push_str(&format!(
360                "  Throughput: {:.2} steps/sec\n",
361                metrics.throughput()
362            ));
363            report.push_str(&format!(
364                "  Learning rate: {:.6}\n",
365                metrics.current_learning_rate
366            ));
367            report.push_str(&format!(
368                "  Gradient norm: {:.6}\n",
369                metrics.gradient_stats.norm
370            ));
371            report.push_str(&format!(
372                "  Update magnitude: {:.6}\n",
373                metrics.parameter_stats.update_magnitude
374            ));
375            report.push_str(&format!(
376                "  Converging: {}\n",
377                metrics.convergence.is_converging
378            ));
379            report.push_str(&format!(
380                "  Memory usage: {} bytes\n\n",
381                metrics.memory_usage
382            ));
383        }
384
385        report
386    }
387}
388
389impl Default for MetricsCollector {
390    fn default() -> Self {
391        Self::new()
392    }
393}
394
395/// Metrics reporter for exporting metrics to various formats
396pub struct MetricsReporter;
397
398impl MetricsReporter {
399    /// Export metrics to JSON format
400    pub fn to_json(metrics: &OptimizerMetrics) -> String {
401        format!(
402            r#"{{
403  "name": "{}",
404  "step_count": {},
405  "avg_step_time_ms": {},
406  "throughput": {},
407  "learning_rate": {},
408  "gradient_norm": {},
409  "update_magnitude": {},
410  "is_converging": {}
411}}"#,
412            metrics.name,
413            metrics.step_count,
414            metrics.avg_step_time.as_millis(),
415            metrics.throughput(),
416            metrics.current_learning_rate,
417            metrics.gradient_stats.norm,
418            metrics.parameter_stats.update_magnitude,
419            metrics.convergence.is_converging
420        )
421    }
422
423    /// Export metrics to CSV format
424    pub fn to_csv_header() -> String {
425        "name,step_count,avg_step_time_ms,throughput,learning_rate,gradient_norm,update_magnitude,is_converging".to_string()
426    }
427
428    /// Export metrics to CSV row
429    pub fn to_csv(metrics: &OptimizerMetrics) -> String {
430        format!(
431            "{},{},{},{},{},{},{},{}",
432            metrics.name,
433            metrics.step_count,
434            metrics.avg_step_time.as_millis(),
435            metrics.throughput(),
436            metrics.current_learning_rate,
437            metrics.gradient_stats.norm,
438            metrics.parameter_stats.update_magnitude,
439            metrics.convergence.is_converging
440        )
441    }
442}
443
444#[cfg(test)]
445mod tests {
446    use super::*;
447    use scirs2_core::ndarray::Array1;
448
449    #[test]
450    fn test_optimizer_metrics_creation() {
451        let metrics = OptimizerMetrics::new("sgd");
452        assert_eq!(metrics.name, "sgd");
453        assert_eq!(metrics.step_count, 0);
454        assert_eq!(metrics.throughput(), 0.0);
455    }
456
457    #[test]
458    fn test_gradient_statistics() {
459        let mut stats = GradientStatistics::default();
460        let grads = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0]);
461        stats
462            .update(&grads.view())
463            .expect("f64 gradients are representable");
464
465        assert!((stats.mean - 3.0).abs() < 1e-6);
466        assert!(stats.max > 4.9);
467        assert!(stats.min < 1.1);
468        assert!(stats.norm > 0.0);
469    }
470
471    #[test]
472    fn test_parameter_statistics() {
473        let mut stats = ParameterStatistics::default();
474        let before = Array1::from_vec(vec![1.0, 2.0, 3.0]);
475        let after = Array1::from_vec(vec![0.9, 1.9, 2.9]);
476        stats
477            .update(&before.view(), &after.view())
478            .expect("f64 parameters of equal length are representable");
479
480        assert!(stats.update_magnitude > 0.0);
481        assert!(stats.relative_change > 0.0);
482        assert!((stats.mean - 1.9).abs() < 1e-6);
483    }
484
485    #[test]
486    fn test_metrics_collector() {
487        let mut collector = MetricsCollector::new();
488        collector.register_optimizer("sgd");
489
490        let grads = Array1::from_vec(vec![0.1, 0.2, 0.3]);
491        let before = Array1::from_vec(vec![1.0, 2.0, 3.0]);
492        let after = Array1::from_vec(vec![0.99, 1.98, 2.97]);
493
494        let result = collector.update(
495            "sgd",
496            Duration::from_millis(10),
497            0.01,
498            &grads.view(),
499            &before.view(),
500            &after.view(),
501        );
502
503        assert!(result.is_ok());
504        let metrics = collector.get_metrics("sgd").expect("unwrap failed");
505        assert_eq!(metrics.step_count, 1);
506    }
507
508    #[test]
509    fn test_metrics_collector_multiple_updates() {
510        let mut collector = MetricsCollector::new();
511        collector.register_optimizer("adam");
512
513        let grads = Array1::from_vec(vec![0.1, 0.2]);
514        let before = Array1::from_vec(vec![1.0, 2.0]);
515        let after = Array1::from_vec(vec![0.99, 1.98]);
516
517        for _ in 0..10 {
518            collector
519                .update(
520                    "adam",
521                    Duration::from_millis(5),
522                    0.001,
523                    &grads.view(),
524                    &before.view(),
525                    &after.view(),
526                )
527                .expect("unwrap failed");
528        }
529
530        let metrics = collector.get_metrics("adam").expect("unwrap failed");
531        assert_eq!(metrics.step_count, 10);
532        assert!(metrics.throughput() > 0.0);
533    }
534
535    #[test]
536    fn test_metrics_reset() {
537        let mut metrics = OptimizerMetrics::new("test");
538        let grads = Array1::from_vec(vec![0.1]);
539        let before = Array1::from_vec(vec![1.0]);
540        let after = Array1::from_vec(vec![0.99]);
541
542        metrics
543            .update_step(
544                Duration::from_millis(10),
545                0.01,
546                &grads.view(),
547                &before.view(),
548                &after.view(),
549            )
550            .expect("well-formed f64 step must record");
551
552        assert_eq!(metrics.step_count, 1);
553
554        metrics.reset();
555        assert_eq!(metrics.step_count, 0);
556        assert_eq!(metrics.total_step_time, Duration::ZERO);
557    }
558
559    #[test]
560    fn test_summary_report() {
561        let mut collector = MetricsCollector::new();
562        collector.register_optimizer("sgd");
563
564        let grads = Array1::from_vec(vec![0.1]);
565        let before = Array1::from_vec(vec![1.0]);
566        let after = Array1::from_vec(vec![0.99]);
567
568        collector
569            .update(
570                "sgd",
571                Duration::from_millis(10),
572                0.01,
573                &grads.view(),
574                &before.view(),
575                &after.view(),
576            )
577            .expect("unwrap failed");
578
579        let report = collector.summary_report();
580        assert!(report.contains("Optimizer: sgd"));
581        assert!(report.contains("Steps: 1"));
582    }
583
584    #[test]
585    fn test_metrics_reporter_json() {
586        let metrics = OptimizerMetrics::new("test");
587        let json = MetricsReporter::to_json(&metrics);
588        assert!(json.contains("\"name\": \"test\""));
589        assert!(json.contains("\"step_count\": 0"));
590    }
591
592    #[test]
593    fn test_metrics_reporter_csv() {
594        let metrics = OptimizerMetrics::new("test");
595        let header = MetricsReporter::to_csv_header();
596        let row = MetricsReporter::to_csv(&metrics);
597
598        assert!(header.contains("name"));
599        assert!(header.contains("step_count"));
600        assert!(row.starts_with("test,0,"));
601    }
602
603    #[test]
604    fn test_convergence_metrics() {
605        let mut convergence = ConvergenceMetrics::default();
606
607        // Update with some values
608        let mut param_stats = ParameterStatistics {
609            update_magnitude: 1.0,
610            ..Default::default()
611        };
612        convergence.update(&param_stats);
613        assert_eq!(convergence.update_moving_avg, 0.1);
614
615        param_stats.update_magnitude = 0.5;
616        convergence.update(&param_stats);
617        // update_moving_avg = 0.1 * 0.5 + 0.9 * 0.1 = 0.14
618        assert!((convergence.update_moving_avg - 0.14).abs() < 1e-6);
619
620        // Verify convergence detection works
621        param_stats.update_magnitude = 0.05;
622        convergence.update(&param_stats);
623        // Should detect converging since 0.05 < 0.14
624        assert!(convergence.is_converging);
625        assert!(convergence.update_moving_avg > 0.0);
626    }
627}