Skip to main content

torsh_optim/benchmarks/
optimizer.rs

1//! Core optimizer benchmarking implementation
2//!
3//! This module contains the main OptimizerBenchmarks struct and its implementation
4//! for comprehensive optimizer performance evaluation.
5
6use super::core::{BenchmarkConfig, BenchmarkResult, MemoryStats};
7use crate::{Optimizer, OptimizerResult};
8use std::time::{Duration, Instant};
9use torsh_core::device::DeviceType;
10use torsh_tensor::{creation, Tensor};
11
12/// Optimizer benchmark suite
13pub struct OptimizerBenchmarks {
14    config: BenchmarkConfig,
15}
16
17impl OptimizerBenchmarks {
18    /// Create a new benchmark suite
19    pub fn new() -> Self {
20        Self {
21            config: BenchmarkConfig::default(),
22        }
23    }
24
25    /// Create a new benchmark suite with custom configuration
26    pub fn with_config(config: BenchmarkConfig) -> Self {
27        Self { config }
28    }
29
30    /// Estimate memory usage for parameters and optimizer state
31    fn estimate_memory_usage<O: Optimizer>(params: &[Tensor], optimizer: &O) -> usize {
32        let mut total_bytes = 0;
33
34        // Calculate memory for parameters
35        for param in params {
36            let shape = param.shape();
37            let element_count = shape.dims().iter().product::<usize>();
38            // Assume f32 elements (4 bytes each)
39            total_bytes += element_count * 4;
40
41            // Add memory for gradients if present
42            if param.has_grad() {
43                total_bytes += element_count * 4;
44            }
45        }
46
47        // Estimate optimizer state memory
48        // This is a rough approximation based on common optimizer patterns
49        let state_multiplier = match optimizer.get_lr().len() {
50            // Simple optimizers like SGD might have minimal state
51            1 => 1.2,
52            // More complex optimizers like Adam have momentum and squared gradients
53            _ => 3.0,
54        };
55
56        // Base memory for optimizer structure itself (at least 1KB)
57        let base_optimizer_memory = 1024;
58        let optimizer_state_bytes = if total_bytes > 0 {
59            (total_bytes as f64 * state_multiplier) as usize
60        } else {
61            base_optimizer_memory
62        };
63
64        total_bytes + optimizer_state_bytes
65    }
66
67    /// Benchmark optimizer step performance
68    pub fn benchmark_step_performance<O: Optimizer>(
69        &self,
70        mut optimizer: O,
71        problem_size: usize,
72    ) -> OptimizerResult<BenchmarkResult> {
73        let mut params = creation::randn::<f32>(&[problem_size])?;
74        let mut iteration_times = Vec::new();
75
76        // Warmup
77        for _ in 0..self.config.warmup_iterations {
78            let grads = creation::randn::<f32>(&[problem_size])?;
79            params.set_grad(Some(grads));
80            optimizer.step()?;
81        }
82
83        let start_time = Instant::now();
84        let mut iterations_completed = 0;
85
86        // Benchmark loop
87        for i in 0..self.config.num_iterations {
88            // Check time limit
89            if start_time.elapsed().as_secs_f32() > self.config.max_time_seconds {
90                break;
91            }
92
93            let grads = creation::randn::<f32>(&[problem_size])?;
94            params.set_grad(Some(grads));
95
96            let iter_start = Instant::now();
97            optimizer.step()?;
98            let iter_time = iter_start.elapsed();
99
100            iteration_times.push(iter_time);
101            iterations_completed += 1;
102        }
103
104        let total_time = start_time.elapsed();
105
106        // Calculate statistics
107        let avg_time = total_time / iterations_completed as u32;
108        let min_time = iteration_times
109            .iter()
110            .min()
111            .copied()
112            .unwrap_or(Duration::ZERO);
113        let max_time = iteration_times
114            .iter()
115            .max()
116            .copied()
117            .unwrap_or(Duration::ZERO);
118
119        // Calculate standard deviation
120        let mean_nanos = avg_time.as_nanos() as f64;
121        let variance = iteration_times
122            .iter()
123            .map(|t| (t.as_nanos() as f64 - mean_nanos).powi(2))
124            .sum::<f64>()
125            / iterations_completed as f64;
126        let std_dev = Duration::from_nanos(variance.sqrt() as u64);
127
128        Ok(BenchmarkResult {
129            name: format!("step_performance_size_{}", problem_size),
130            iterations_completed,
131            total_time,
132            avg_time_per_iteration: avg_time,
133            min_time_per_iteration: min_time,
134            max_time_per_iteration: max_time,
135            time_std_dev: std_dev,
136            final_loss: None,
137            memory_stats: None,
138            convergence_rate: None,
139        })
140    }
141
142    /// Benchmark convergence on quadratic function
143    pub fn benchmark_quadratic_convergence<O: Optimizer>(
144        &self,
145        mut optimizer: O,
146        dimension: usize,
147    ) -> OptimizerResult<BenchmarkResult> {
148        let device = self.config.device;
149        let mut params = creation::randn::<f32>(&[dimension])?;
150        let target = Tensor::zeros(&[dimension], DeviceType::Cpu)?;
151
152        let mut losses = Vec::new();
153        let mut iteration_times = Vec::new();
154
155        // Initial loss
156        let initial_loss = params.sub(&target)?.pow(2.0)?.sum()?.item()?;
157        losses.push(initial_loss);
158
159        let start_time = Instant::now();
160        let mut iterations_completed = 0;
161
162        for i in 0..self.config.num_iterations {
163            // Check time limit
164            if start_time.elapsed().as_secs_f32() > self.config.max_time_seconds {
165                break;
166            }
167
168            // Compute gradients for quadratic loss: grad = 2 * (params - target)
169            let grads = params.sub(&target)?.mul_scalar(2.0)?;
170
171            let iter_start = Instant::now();
172            params.set_grad(Some(grads));
173            optimizer.step()?;
174            let iter_time = iter_start.elapsed();
175
176            iteration_times.push(iter_time);
177
178            // Compute loss
179            let loss = params.sub(&target)?.pow(2.0)?.sum()?.item()?;
180            losses.push(loss);
181
182            iterations_completed += 1;
183
184            // Early stopping if converged
185            if loss < 1e-8 {
186                break;
187            }
188        }
189
190        let total_time = start_time.elapsed();
191        let final_loss = losses.last().copied().unwrap_or(f32::INFINITY);
192
193        // Calculate convergence rate
194        let convergence_rate = if losses.len() > 1 {
195            let log_reduction = (initial_loss.ln() - final_loss.ln()).max(0.0);
196            Some(log_reduction / iterations_completed as f32)
197        } else {
198            None
199        };
200
201        // Calculate timing statistics
202        let avg_time = total_time / iterations_completed as u32;
203        let min_time = iteration_times
204            .iter()
205            .min()
206            .copied()
207            .unwrap_or(Duration::ZERO);
208        let max_time = iteration_times
209            .iter()
210            .max()
211            .copied()
212            .unwrap_or(Duration::ZERO);
213
214        let mean_nanos = avg_time.as_nanos() as f64;
215        let variance = iteration_times
216            .iter()
217            .map(|t| (t.as_nanos() as f64 - mean_nanos).powi(2))
218            .sum::<f64>()
219            / iterations_completed as f64;
220        let std_dev = Duration::from_nanos(variance.sqrt() as u64);
221
222        Ok(BenchmarkResult {
223            name: format!("quadratic_convergence_dim_{}", dimension),
224            iterations_completed,
225            total_time,
226            avg_time_per_iteration: avg_time,
227            min_time_per_iteration: min_time,
228            max_time_per_iteration: max_time,
229            time_std_dev: std_dev,
230            final_loss: Some(final_loss),
231            memory_stats: None,
232            convergence_rate,
233        })
234    }
235
236    /// Benchmark memory usage scaling
237    pub fn benchmark_memory_scaling<O: Optimizer + Clone>(
238        &self,
239        optimizer_factory: impl Fn(Vec<Tensor>) -> OptimizerResult<O>,
240    ) -> OptimizerResult<Vec<BenchmarkResult>> {
241        let device = self.config.device;
242        let sizes = vec![100, 1000, 10000, 100000];
243        let mut results = Vec::new();
244
245        for &size in &sizes {
246            let params = vec![creation::randn::<f32>(&[size])?];
247            let mut optimizer = optimizer_factory(params.clone())?;
248
249            // Warmup
250            for _ in 0..10 {
251                let grads = creation::randn::<f32>(&[size])?;
252                params[0].set_grad(Some(grads));
253                optimizer.step()?;
254            }
255
256            // Measure memory usage during optimization
257            let start_time = Instant::now();
258            let mut iteration_times = Vec::new();
259            let iterations = 100.min(self.config.num_iterations);
260
261            // Track memory usage if enabled
262            let memory_stats = if self.config.profile_memory {
263                let initial_memory = Self::estimate_memory_usage(&params, &optimizer);
264                let mut memory_samples = Vec::new();
265                let mut peak_memory = initial_memory;
266
267                memory_samples.push(initial_memory);
268
269                for _ in 0..iterations {
270                    let grads = creation::randn::<f32>(&[size])?;
271
272                    let iter_start = Instant::now();
273                    params[0].set_grad(Some(grads));
274                    optimizer.step()?;
275                    let iter_time = iter_start.elapsed();
276
277                    iteration_times.push(iter_time);
278
279                    // Sample memory usage periodically (every 10 iterations)
280                    if iteration_times.len() % 10 == 0 {
281                        let current_memory = Self::estimate_memory_usage(&params, &optimizer);
282                        memory_samples.push(current_memory);
283                        peak_memory = peak_memory.max(current_memory);
284                    }
285                }
286
287                let final_memory = Self::estimate_memory_usage(&params, &optimizer);
288                memory_samples.push(final_memory);
289                peak_memory = peak_memory.max(final_memory);
290
291                let avg_memory = memory_samples.iter().sum::<usize>() / memory_samples.len().max(1);
292
293                Some(MemoryStats {
294                    peak_memory_bytes: peak_memory,
295                    initial_memory_bytes: initial_memory,
296                    final_memory_bytes: final_memory,
297                    avg_memory_bytes: avg_memory,
298                })
299            } else {
300                // If memory profiling is disabled, still collect timing data
301                for _ in 0..iterations {
302                    let grads = creation::randn::<f32>(&[size])?;
303
304                    let iter_start = Instant::now();
305                    params[0].set_grad(Some(grads));
306                    optimizer.step()?;
307                    let iter_time = iter_start.elapsed();
308
309                    iteration_times.push(iter_time);
310                }
311
312                None
313            };
314
315            let total_time = start_time.elapsed();
316            let avg_time = total_time / iterations as u32;
317            let min_time = iteration_times
318                .iter()
319                .min()
320                .copied()
321                .unwrap_or(Duration::ZERO);
322            let max_time = iteration_times
323                .iter()
324                .max()
325                .copied()
326                .unwrap_or(Duration::ZERO);
327
328            let mean_nanos = avg_time.as_nanos() as f64;
329            let variance = iteration_times
330                .iter()
331                .map(|t| (t.as_nanos() as f64 - mean_nanos).powi(2))
332                .sum::<f64>()
333                / iterations as f64;
334            let std_dev = Duration::from_nanos(variance.sqrt() as u64);
335
336            results.push(BenchmarkResult {
337                name: format!("memory_scaling_size_{}", size),
338                iterations_completed: iterations,
339                total_time,
340                avg_time_per_iteration: avg_time,
341                min_time_per_iteration: min_time,
342                max_time_per_iteration: max_time,
343                time_std_dev: std_dev,
344                final_loss: None,
345                memory_stats,
346                convergence_rate: None,
347            });
348        }
349
350        Ok(results)
351    }
352
353    /// Benchmark sparse gradient handling
354    pub fn benchmark_sparse_gradients<O: Optimizer>(
355        &self,
356        mut optimizer: O,
357        total_params: usize,
358        sparsity: f32,
359    ) -> OptimizerResult<BenchmarkResult> {
360        let mut params = creation::randn::<f32>(&[total_params])?;
361
362        let mut iteration_times = Vec::new();
363
364        // Warmup
365        for _ in 0..self.config.warmup_iterations {
366            let mut grads = Tensor::zeros(&[total_params], DeviceType::Cpu)?;
367
368            // Set sparse gradients
369            for i in 0..total_params {
370                if (i as f32 / total_params as f32) < sparsity {
371                    let grad_val = ((i as f32 * 0.1) % 2.0) - 1.0;
372                    let _ = grads.set(&[i], grad_val);
373                }
374            }
375
376            params.set_grad(Some(grads));
377            optimizer.step()?;
378        }
379
380        let start_time = Instant::now();
381        let mut iterations_completed = 0;
382
383        for _ in 0..self.config.num_iterations {
384            if start_time.elapsed().as_secs_f32() > self.config.max_time_seconds {
385                break;
386            }
387
388            let mut grads = Tensor::zeros(&[total_params], DeviceType::Cpu)?;
389
390            // Set sparse gradients
391            for i in 0..total_params {
392                if (i as f32 / total_params as f32) < sparsity {
393                    let grad_val = ((i as f32 * 0.1) % 2.0) - 1.0;
394                    let _ = grads.set(&[i], grad_val);
395                }
396            }
397
398            let iter_start = Instant::now();
399            params.set_grad(Some(grads));
400            optimizer.step()?;
401            let iter_time = iter_start.elapsed();
402
403            iteration_times.push(iter_time);
404            iterations_completed += 1;
405        }
406
407        let total_time = start_time.elapsed();
408        let avg_time = total_time / iterations_completed as u32;
409        let min_time = iteration_times
410            .iter()
411            .min()
412            .copied()
413            .unwrap_or(Duration::ZERO);
414        let max_time = iteration_times
415            .iter()
416            .max()
417            .copied()
418            .unwrap_or(Duration::ZERO);
419
420        let mean_nanos = avg_time.as_nanos() as f64;
421        let variance = iteration_times
422            .iter()
423            .map(|t| (t.as_nanos() as f64 - mean_nanos).powi(2))
424            .sum::<f64>()
425            / iterations_completed as f64;
426        let std_dev = Duration::from_nanos(variance.sqrt() as u64);
427
428        Ok(BenchmarkResult {
429            name: format!(
430                "sparse_gradients_params_{}_sparsity_{:.2}",
431                total_params, sparsity
432            ),
433            iterations_completed,
434            total_time,
435            avg_time_per_iteration: avg_time,
436            min_time_per_iteration: min_time,
437            max_time_per_iteration: max_time,
438            time_std_dev: std_dev,
439            final_loss: None,
440            memory_stats: None,
441            convergence_rate: None,
442        })
443    }
444
445    /// Run comprehensive benchmark suite
446    pub fn run_comprehensive_benchmarks<O: Optimizer + Clone>(
447        &self,
448        optimizer: O,
449    ) -> OptimizerResult<Vec<BenchmarkResult>> {
450        let mut results = Vec::new();
451
452        // Step performance benchmarks for different problem sizes
453        for &size in &[100, 1000, 10000] {
454            results.push(self.benchmark_step_performance(optimizer.clone(), size)?);
455        }
456
457        // Convergence benchmarks
458        for &dim in &[10, 100, 1000] {
459            results.push(self.benchmark_quadratic_convergence(optimizer.clone(), dim)?);
460        }
461
462        // Sparse gradient benchmarks
463        for &sparsity in &[0.1, 0.01, 0.001] {
464            results.push(self.benchmark_sparse_gradients(optimizer.clone(), 10000, sparsity)?);
465        }
466
467        Ok(results)
468    }
469
470    /// Print benchmark results in a formatted table
471    pub fn print_results(&self, results: &[BenchmarkResult]) {
472        println!("\n{:=<100}", "");
473        println!("{:^100}", "OPTIMIZER BENCHMARK RESULTS");
474        println!("{:=<100}", "");
475
476        println!(
477            "{:<40} {:>12} {:>12} {:>12} {:>12} {:>10}",
478            "Benchmark", "Iterations", "Total Time", "Avg Time", "Min Time", "Max Time"
479        );
480        println!("{:-<100}", "");
481
482        for result in results {
483            println!(
484                "{:<40} {:>12} {:>12.3?} {:>12.3?} {:>12.3?} {:>12.3?}",
485                result.name,
486                result.iterations_completed,
487                result.total_time,
488                result.avg_time_per_iteration,
489                result.min_time_per_iteration,
490                result.max_time_per_iteration
491            );
492
493            if let Some(loss) = result.final_loss {
494                println!("{:<40} Final Loss: {:.6e}", "", loss);
495            }
496
497            if let Some(rate) = result.convergence_rate {
498                println!("{:<40} Convergence Rate: {:.6e}", "", rate);
499            }
500        }
501
502        println!("{:=<100}", "");
503    }
504
505    /// Benchmark noisy quadratic convergence (more realistic optimization scenario)
506    pub fn benchmark_noisy_quadratic_convergence<O: Optimizer>(
507        &self,
508        mut optimizer: O,
509        dimension: usize,
510        noise_level: f32,
511    ) -> OptimizerResult<BenchmarkResult> {
512        let mut params = creation::randn::<f32>(&[dimension])?;
513        let target = Tensor::zeros(&[dimension], DeviceType::Cpu)?;
514
515        let mut losses = Vec::new();
516        let mut iteration_times = Vec::new();
517
518        // Initial loss
519        let initial_loss = params.sub(&target)?.pow(2.0)?.sum()?.item()?;
520        losses.push(initial_loss);
521
522        let start_time = Instant::now();
523        let mut iterations_completed = 0;
524
525        for _ in 0..self.config.num_iterations {
526            if start_time.elapsed().as_secs_f32() > self.config.max_time_seconds {
527                break;
528            }
529
530            // Compute gradients for quadratic loss with noise: grad = 2 * (params - target) + noise
531            let clean_grads = params.sub(&target)?.mul_scalar(2.0)?;
532            let noise = creation::randn::<f32>(&[dimension])?.mul_scalar(noise_level)?;
533            let grads = clean_grads.add(&noise)?;
534
535            let iter_start = Instant::now();
536            params.set_grad(Some(grads));
537            optimizer.step()?;
538            let iter_time = iter_start.elapsed();
539
540            iteration_times.push(iter_time);
541
542            // Compute loss
543            let loss = params.sub(&target)?.pow(2.0)?.sum()?.item()?;
544            losses.push(loss);
545
546            iterations_completed += 1;
547
548            // Early stopping if converged (more lenient due to noise)
549            if loss < 1e-6 {
550                break;
551            }
552        }
553
554        let total_time = start_time.elapsed();
555        let final_loss = losses.last().copied().unwrap_or(f32::INFINITY);
556
557        // Calculate convergence rate
558        let convergence_rate = if losses.len() > 1 {
559            let log_reduction = (initial_loss.ln() - final_loss.ln()).max(0.0);
560            Some(log_reduction / iterations_completed as f32)
561        } else {
562            None
563        };
564
565        // Calculate timing statistics
566        let avg_time = total_time / iterations_completed as u32;
567        let min_time = iteration_times
568            .iter()
569            .min()
570            .copied()
571            .unwrap_or(Duration::ZERO);
572        let max_time = iteration_times
573            .iter()
574            .max()
575            .copied()
576            .unwrap_or(Duration::ZERO);
577
578        let mean_nanos = avg_time.as_nanos() as f64;
579        let variance = iteration_times
580            .iter()
581            .map(|t| (t.as_nanos() as f64 - mean_nanos).powi(2))
582            .sum::<f64>()
583            / iterations_completed as f64;
584        let std_dev = Duration::from_nanos(variance.sqrt() as u64);
585
586        Ok(BenchmarkResult {
587            name: format!("noisy_quadratic_dim_{}_noise_{:.3}", dimension, noise_level),
588            iterations_completed,
589            total_time,
590            avg_time_per_iteration: avg_time,
591            min_time_per_iteration: min_time,
592            max_time_per_iteration: max_time,
593            time_std_dev: std_dev,
594            final_loss: Some(final_loss),
595            memory_stats: None,
596            convergence_rate,
597        })
598    }
599
600    /// Benchmark Rosenbrock function optimization (classic non-convex benchmark)
601    pub fn benchmark_rosenbrock_optimization<O: Optimizer>(
602        &self,
603        mut optimizer: O,
604        dimension: usize,
605    ) -> OptimizerResult<BenchmarkResult> {
606        let mut params = creation::randn::<f32>(&[dimension])?;
607
608        let mut losses = Vec::new();
609        let mut iteration_times = Vec::new();
610
611        let start_time = Instant::now();
612        let mut iterations_completed = 0;
613
614        for _ in 0..self.config.num_iterations {
615            if start_time.elapsed().as_secs_f32() > self.config.max_time_seconds {
616                break;
617            }
618
619            // Compute Rosenbrock function gradients
620            // f(x) = sum(100*(x[i+1] - x[i]^2)^2 + (1 - x[i])^2) for i = 0..n-2
621            let mut grads = Tensor::zeros(&[dimension], DeviceType::Cpu)?;
622
623            // This is a simplified approximation of Rosenbrock gradients
624            // for the purpose of benchmarking optimizer performance
625            for i in 0..dimension {
626                let x_i = params.get(&[i])?;
627                let grad_val = if i < dimension - 1 {
628                    let x_next = params.get(&[i + 1])?;
629                    // Approximate gradient: -400*x[i]*(x[i+1] - x[i]^2) - 2*(1 - x[i])
630                    -400.0 * x_i * (x_next - x_i * x_i) - 2.0 * (1.0 - x_i)
631                } else {
632                    // Last element: 200*(x[i] - x[i-1]^2)
633                    if i > 0 {
634                        let x_prev = params.get(&[i - 1])?;
635                        200.0 * (x_i - x_prev * x_prev)
636                    } else {
637                        0.0
638                    }
639                };
640                let _ = grads.set(&[i], grad_val);
641            }
642
643            let iter_start = Instant::now();
644            params.set_grad(Some(grads));
645            optimizer.step()?;
646            let iter_time = iter_start.elapsed();
647
648            iteration_times.push(iter_time);
649
650            // Compute approximate Rosenbrock loss
651            let mut loss = 0.0;
652            for i in 0..(dimension - 1) {
653                let x_i = params.get(&[i])?;
654                let x_next = params.get(&[i + 1])?;
655                loss += 100.0 * (x_next - x_i * x_i).powf(2.0) + (1.0 - x_i).powf(2.0);
656            }
657            losses.push(loss);
658
659            iterations_completed += 1;
660
661            // Early stopping if converged
662            if loss < 1e-4 {
663                break;
664            }
665        }
666
667        let total_time = start_time.elapsed();
668        let final_loss = losses.last().copied().unwrap_or(f32::INFINITY);
669        let initial_loss = losses.first().copied().unwrap_or(0.0);
670
671        // Calculate convergence rate
672        let convergence_rate = if losses.len() > 1 && initial_loss > 0.0 {
673            let log_reduction = (initial_loss.ln() - final_loss.ln()).max(0.0);
674            Some(log_reduction / iterations_completed as f32)
675        } else {
676            None
677        };
678
679        // Calculate timing statistics
680        let avg_time = total_time / iterations_completed as u32;
681        let min_time = iteration_times
682            .iter()
683            .min()
684            .copied()
685            .unwrap_or(Duration::ZERO);
686        let max_time = iteration_times
687            .iter()
688            .max()
689            .copied()
690            .unwrap_or(Duration::ZERO);
691
692        let mean_nanos = avg_time.as_nanos() as f64;
693        let variance = iteration_times
694            .iter()
695            .map(|t| (t.as_nanos() as f64 - mean_nanos).powi(2))
696            .sum::<f64>()
697            / iterations_completed as f64;
698        let std_dev = Duration::from_nanos(variance.sqrt() as u64);
699
700        Ok(BenchmarkResult {
701            name: format!("rosenbrock_optimization_dim_{}", dimension),
702            iterations_completed,
703            total_time,
704            avg_time_per_iteration: avg_time,
705            min_time_per_iteration: min_time,
706            max_time_per_iteration: max_time,
707            time_std_dev: std_dev,
708            final_loss: Some(final_loss),
709            memory_stats: None,
710            convergence_rate,
711        })
712    }
713}
714
715impl Default for OptimizerBenchmarks {
716    fn default() -> Self {
717        Self::new()
718    }
719}
720
721#[cfg(test)]
722mod tests {
723    use super::*;
724    use crate::sgd::SGD;
725    use parking_lot::RwLock;
726    use std::sync::Arc;
727
728    #[test]
729    fn test_optimizer_benchmarks_creation() {
730        let benchmarks = OptimizerBenchmarks::new();
731        assert_eq!(benchmarks.config.num_iterations, 1000);
732        assert_eq!(benchmarks.config.warmup_iterations, 100);
733        assert_eq!(benchmarks.config.max_time_seconds, 60.0);
734    }
735
736    #[test]
737    fn test_custom_config() {
738        let custom_config = BenchmarkConfig {
739            num_iterations: 500,
740            warmup_iterations: 50,
741            max_time_seconds: 30.0,
742            ..Default::default()
743        };
744
745        let benchmarks = OptimizerBenchmarks::with_config(custom_config);
746        assert_eq!(benchmarks.config.num_iterations, 500);
747        assert_eq!(benchmarks.config.warmup_iterations, 50);
748        assert_eq!(benchmarks.config.max_time_seconds, 30.0);
749    }
750
751    #[test]
752    fn test_memory_estimation() {
753        // This is a basic test - in practice we'd need actual tensors and optimizers
754        // For now, just ensure the function doesn't panic
755        let params: Vec<Tensor> = Vec::new();
756        let sgd_params = vec![Arc::new(RwLock::new(
757            creation::zeros::<f32>(&[10]).expect("operation should succeed"),
758        ))];
759        let sgd = SGD::new(sgd_params, 0.01, None, None, None, false);
760
761        let memory = OptimizerBenchmarks::estimate_memory_usage(&params, &sgd);
762        assert!(memory > 0); // Should account for optimizer state even with no params
763    }
764}