Skip to main content

torsh_optim/benchmarks/
comparison.rs

1//! Multi-optimizer comparison functionality
2//!
3//! This module provides comprehensive comparison capabilities for evaluating
4//! multiple optimizers side-by-side with statistical analysis.
5
6use super::core::{BenchmarkConfig, BenchmarkResult, MemoryStats, StatisticalAnalysis};
7use crate::{Optimizer, OptimizerResult};
8use std::time::Duration;
9
10/// Comparison results between multiple optimizers
11#[derive(Debug, Clone)]
12#[cfg_attr(feature = "serialize", derive(serde::Serialize))]
13pub struct ComparisonResult {
14    /// Optimizer name
15    pub optimizer_name: String,
16    /// Individual benchmark results
17    pub benchmark_results: Vec<BenchmarkResult>,
18    /// Statistical analysis
19    pub statistical_analysis: StatisticalAnalysis,
20    /// Performance ranking (1 = best)
21    pub performance_rank: usize,
22    /// Convergence ranking (1 = best)
23    pub convergence_rank: usize,
24    /// Memory efficiency ranking (1 = best)
25    pub memory_rank: usize,
26}
27
28/// Multi-optimizer comparison suite
29pub struct OptimizerComparison<O: Optimizer + Clone> {
30    optimizers: Vec<(String, Box<dyn Fn() -> OptimizerResult<O>>)>,
31    config: BenchmarkConfig,
32}
33
34impl<O: Optimizer + Clone> OptimizerComparison<O> {
35    /// Create a new optimizer comparison suite
36    pub fn new() -> Self {
37        Self {
38            optimizers: Vec::new(),
39            config: BenchmarkConfig::default(),
40        }
41    }
42
43    /// Add an optimizer to the comparison
44    pub fn add_optimizer<F>(mut self, name: &str, factory: F) -> Self
45    where
46        F: Fn() -> OptimizerResult<O> + 'static,
47    {
48        self.optimizers.push((name.to_string(), Box::new(factory)));
49        self
50    }
51
52    /// Set custom benchmark configuration
53    pub fn with_config(mut self, config: BenchmarkConfig) -> Self {
54        self.config = config;
55        self
56    }
57
58    /// Run comprehensive comparison between all optimizers
59    pub fn run_comparison_suite(&self) -> OptimizerResult<Vec<ComparisonResult>> {
60        use super::optimizer::OptimizerBenchmarks;
61
62        let benchmarks = OptimizerBenchmarks::with_config(self.config.clone());
63        let mut comparison_results = Vec::new();
64
65        for (name, factory) in &self.optimizers {
66            let optimizer = factory()?;
67            let results = benchmarks.run_comprehensive_benchmarks(optimizer)?;
68
69            // Calculate statistical analysis
70            let execution_times: Vec<Duration> =
71                results.iter().map(|r| r.avg_time_per_iteration).collect();
72
73            let stats = Self::calculate_statistical_analysis(&execution_times);
74
75            comparison_results.push(ComparisonResult {
76                optimizer_name: name.clone(),
77                benchmark_results: results,
78                statistical_analysis: stats,
79                performance_rank: 0, // Will be calculated after all results are collected
80                convergence_rank: 0,
81                memory_rank: 0,
82            });
83        }
84
85        // Calculate rankings
86        self.calculate_rankings(&mut comparison_results);
87
88        Ok(comparison_results)
89    }
90
91    /// Calculate statistical analysis from timing data
92    fn calculate_statistical_analysis(times: &[Duration]) -> StatisticalAnalysis {
93        if times.is_empty() {
94            return StatisticalAnalysis {
95                mean_time: Duration::ZERO,
96                median_time: Duration::ZERO,
97                std_dev: Duration::ZERO,
98                confidence_interval: (Duration::ZERO, Duration::ZERO),
99                effect_size: None,
100                p_value: None,
101            };
102        }
103
104        let total_nanos: u64 = times.iter().map(|t| t.as_nanos() as u64).sum();
105        let mean_nanos = total_nanos / times.len() as u64;
106        let mean_time = Duration::from_nanos(mean_nanos);
107
108        let mut sorted_times = times.to_vec();
109        sorted_times.sort();
110        let median_time = sorted_times[sorted_times.len() / 2];
111
112        // Calculate standard deviation
113        let variance: f64 = times
114            .iter()
115            .map(|t| {
116                let diff = t.as_nanos() as f64 - mean_nanos as f64;
117                diff * diff
118            })
119            .sum::<f64>()
120            / times.len() as f64;
121        let std_dev = Duration::from_nanos(variance.sqrt() as u64);
122
123        // Calculate 95% confidence interval (assuming normal distribution)
124        let margin_error = 1.96 * (std_dev.as_nanos() as f64) / (times.len() as f64).sqrt();
125        let lower_bound = Duration::from_nanos((mean_nanos as f64 - margin_error) as u64);
126        let upper_bound = Duration::from_nanos((mean_nanos as f64 + margin_error) as u64);
127
128        StatisticalAnalysis {
129            mean_time,
130            median_time,
131            std_dev,
132            confidence_interval: (lower_bound, upper_bound),
133            effect_size: None,
134            p_value: None,
135        }
136    }
137
138    /// Calculate performance rankings
139    fn calculate_rankings(&self, results: &mut [ComparisonResult]) {
140        // Performance ranking (based on mean execution time - lower is better)
141        let mut perf_indices: Vec<usize> = (0..results.len()).collect();
142        perf_indices.sort_by(|&a, &b| {
143            results[a]
144                .statistical_analysis
145                .mean_time
146                .cmp(&results[b].statistical_analysis.mean_time)
147        });
148        for (rank, &idx) in perf_indices.iter().enumerate() {
149            results[idx].performance_rank = rank + 1;
150        }
151
152        // Convergence ranking (based on convergence rate - higher is better)
153        let mut conv_indices: Vec<usize> = (0..results.len()).collect();
154        conv_indices.sort_by(|&a, &b| {
155            let conv_a = results[a]
156                .benchmark_results
157                .iter()
158                .filter_map(|r| r.convergence_rate)
159                .fold(0.0, |acc, x| acc + x);
160            let conv_b = results[b]
161                .benchmark_results
162                .iter()
163                .filter_map(|r| r.convergence_rate)
164                .fold(0.0, |acc, x| acc + x);
165            conv_b
166                .partial_cmp(&conv_a)
167                .unwrap_or(std::cmp::Ordering::Equal)
168        });
169        for (rank, &idx) in conv_indices.iter().enumerate() {
170            results[idx].convergence_rank = rank + 1;
171        }
172
173        // Memory ranking (based on peak memory usage - lower is better)
174        let mut mem_indices: Vec<usize> = (0..results.len()).collect();
175        mem_indices.sort_by(|&a, &b| {
176            let mem_a = results[a]
177                .benchmark_results
178                .iter()
179                .filter_map(|r| r.memory_stats.as_ref())
180                .map(|m| m.peak_memory_bytes)
181                .fold(0, |acc, x| acc + x);
182            let mem_b = results[b]
183                .benchmark_results
184                .iter()
185                .filter_map(|r| r.memory_stats.as_ref())
186                .map(|m| m.peak_memory_bytes)
187                .fold(0, |acc, x| acc + x);
188            mem_a.cmp(&mem_b)
189        });
190        for (rank, &idx) in mem_indices.iter().enumerate() {
191            results[idx].memory_rank = rank + 1;
192        }
193    }
194
195    /// Print detailed comparison table
196    pub fn print_comparison_table(&self, results: &[ComparisonResult]) {
197        println!("\n{:=<120}", "");
198        println!("{:^120}", "OPTIMIZER COMPARISON RESULTS");
199        println!("{:=<120}", "");
200
201        println!(
202            "{:<20} {:>12} {:>15} {:>15} {:>15} {:>12} {:>12} {:>12}",
203            "Optimizer",
204            "Perf Rank",
205            "Mean Time",
206            "Median Time",
207            "Std Dev",
208            "Conv Rank",
209            "Mem Rank",
210            "Total Score"
211        );
212        println!("{:-<120}", "");
213
214        for result in results {
215            let total_score =
216                result.performance_rank + result.convergence_rank + result.memory_rank;
217            println!(
218                "{:<20} {:>12} {:>15.3?} {:>15.3?} {:>15.3?} {:>12} {:>12} {:>12}",
219                result.optimizer_name,
220                result.performance_rank,
221                result.statistical_analysis.mean_time,
222                result.statistical_analysis.median_time,
223                result.statistical_analysis.std_dev,
224                result.convergence_rank,
225                result.memory_rank,
226                total_score
227            );
228        }
229
230        println!("{:-<120}", "");
231        println!("Lower ranks are better. Total Score = Performance + Convergence + Memory ranks");
232        println!("{:=<120}", "");
233    }
234
235    /// Export comparison results to JSON
236    #[cfg(feature = "serde")]
237    pub fn export_json(&self, results: &[ComparisonResult], filename: &str) -> OptimizerResult<()> {
238        use std::fs::File;
239        use std::io::Write;
240
241        let json = serde_json::to_string_pretty(results)
242            .map_err(|e| crate::OptimizerError::SerializationError(e.to_string()))?;
243
244        let mut file = File::create(filename).map_err(|e| crate::OptimizerError::IoError(e))?;
245
246        file.write_all(json.as_bytes())
247            .map_err(|e| crate::OptimizerError::IoError(e))?;
248
249        println!("Comparison results exported to {}", filename);
250        Ok(())
251    }
252}
253
254impl<O: Optimizer + Clone> Default for OptimizerComparison<O> {
255    fn default() -> Self {
256        Self::new()
257    }
258}
259
260#[cfg(test)]
261mod tests {
262    use super::*;
263    use std::time::Duration;
264
265    #[test]
266    fn test_comparison_result() {
267        let result = ComparisonResult {
268            optimizer_name: "TestOptimizer".to_string(),
269            benchmark_results: Vec::new(),
270            statistical_analysis: StatisticalAnalysis {
271                mean_time: Duration::from_millis(10),
272                median_time: Duration::from_millis(9),
273                std_dev: Duration::from_millis(2),
274                confidence_interval: (Duration::from_millis(8), Duration::from_millis(12)),
275                effect_size: Some(0.5),
276                p_value: Some(0.05),
277            },
278            performance_rank: 1,
279            convergence_rank: 2,
280            memory_rank: 1,
281        };
282
283        assert_eq!(result.optimizer_name, "TestOptimizer");
284        assert_eq!(result.performance_rank, 1);
285        assert_eq!(result.convergence_rank, 2);
286        assert_eq!(result.memory_rank, 1);
287    }
288
289    #[test]
290    fn test_statistical_analysis_empty() {
291        let stats = OptimizerComparison::<crate::sgd::SGD>::calculate_statistical_analysis(&[]);
292        assert_eq!(stats.mean_time, Duration::ZERO);
293        assert_eq!(stats.median_time, Duration::ZERO);
294        assert_eq!(stats.std_dev, Duration::ZERO);
295    }
296
297    #[test]
298    fn test_statistical_analysis_single_value() {
299        let times = vec![Duration::from_millis(10)];
300        let stats = OptimizerComparison::<crate::sgd::SGD>::calculate_statistical_analysis(&times);
301        assert_eq!(stats.mean_time, Duration::from_millis(10));
302        assert_eq!(stats.median_time, Duration::from_millis(10));
303        assert_eq!(stats.std_dev, Duration::ZERO);
304    }
305}