1use 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#[derive(Debug, Clone)]
37pub struct OptimizerMetrics {
38 pub name: String,
40 pub step_count: u64,
42 pub total_step_time: Duration,
44 pub avg_step_time: Duration,
46 pub current_learning_rate: f64,
48 pub gradient_stats: GradientStatistics,
50 pub parameter_stats: ParameterStatistics,
52 pub convergence: ConvergenceMetrics,
54 pub memory_usage: usize,
56}
57
58impl OptimizerMetrics {
59 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 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 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 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 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#[derive(Debug, Clone, Default)]
125pub struct GradientStatistics {
126 pub mean: f64,
128 pub std_dev: f64,
130 pub max: f64,
132 pub min: f64,
134 pub norm: f64,
136 pub num_zeros: usize,
138}
139
140impl GradientStatistics {
141 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 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#[derive(Debug, Clone, Default)]
177pub struct ParameterStatistics {
178 pub mean: f64,
180 pub std_dev: f64,
182 pub update_magnitude: f64,
184 pub relative_change: f64,
186}
187
188impl ParameterStatistics {
189 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 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#[derive(Debug, Clone, Default)]
248pub struct ConvergenceMetrics {
249 pub update_moving_avg: f64,
251 pub is_converging: bool,
253 pub estimated_steps_to_convergence: Option<u64>,
255 pub convergence_rate: f64,
257}
258
259impl ConvergenceMetrics {
260 pub fn update(&mut self, param_stats: &ParameterStatistics) {
262 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 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
275pub struct MetricsCollector {
277 metrics: HashMap<String, OptimizerMetrics>,
279 start_time: Instant,
281}
282
283impl MetricsCollector {
284 pub fn new() -> Self {
286 Self {
287 metrics: HashMap::new(),
288 start_time: Instant::now(),
289 }
290 }
291
292 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 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 pub fn get_metrics(&self, optimizer_name: &str) -> Option<&OptimizerMetrics> {
328 self.metrics.get(optimizer_name)
329 }
330
331 pub fn all_metrics(&self) -> &HashMap<String, OptimizerMetrics> {
333 &self.metrics
334 }
335
336 pub fn elapsed(&self) -> Duration {
338 self.start_time.elapsed()
339 }
340
341 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 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
395pub struct MetricsReporter;
397
398impl MetricsReporter {
399 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 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 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 let mut param_stats = ParameterStatistics {
609 update_magnitude: 1.0,
610 ..Default::default()
611 };
612 convergence.update(¶m_stats);
613 assert_eq!(convergence.update_moving_avg, 0.1);
614
615 param_stats.update_magnitude = 0.5;
616 convergence.update(¶m_stats);
617 assert!((convergence.update_moving_avg - 0.14).abs() < 1e-6);
619
620 param_stats.update_magnitude = 0.05;
622 convergence.update(¶m_stats);
623 assert!(convergence.is_converging);
625 assert!(convergence.update_moving_avg > 0.0);
626 }
627}