1use super::core::{BenchmarkConfig, BenchmarkResult, MemoryStats, StatisticalAnalysis};
7use crate::{Optimizer, OptimizerResult};
8use std::time::Duration;
9
10#[derive(Debug, Clone)]
12#[cfg_attr(feature = "serialize", derive(serde::Serialize))]
13pub struct ComparisonResult {
14 pub optimizer_name: String,
16 pub benchmark_results: Vec<BenchmarkResult>,
18 pub statistical_analysis: StatisticalAnalysis,
20 pub performance_rank: usize,
22 pub convergence_rank: usize,
24 pub memory_rank: usize,
26}
27
28pub 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 pub fn new() -> Self {
37 Self {
38 optimizers: Vec::new(),
39 config: BenchmarkConfig::default(),
40 }
41 }
42
43 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 pub fn with_config(mut self, config: BenchmarkConfig) -> Self {
54 self.config = config;
55 self
56 }
57
58 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 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, convergence_rank: 0,
81 memory_rank: 0,
82 });
83 }
84
85 self.calculate_rankings(&mut comparison_results);
87
88 Ok(comparison_results)
89 }
90
91 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 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 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 fn calculate_rankings(&self, results: &mut [ComparisonResult]) {
140 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 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 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 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 #[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(×);
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}