Skip to main content

torsh_optim/
stress_tests.rs

1//! Stress tests for ToRSh optimizers
2//!
3//! This module provides comprehensive stress testing for optimizers to ensure
4//! robustness under extreme conditions and high load scenarios.
5
6use crate::{OptimizerError, OptimizerResult};
7use parking_lot::RwLock;
8use std::collections::HashMap;
9use std::sync::Arc;
10use std::time::{Duration, Instant};
11use torsh_tensor::{
12    creation::{randn, zeros},
13    Tensor,
14};
15
16#[allow(dead_code)]
17/// Configuration for stress tests
18#[derive(Debug, Clone)]
19pub struct StressTestConfig {
20    /// Number of optimization steps to run
21    pub num_steps: usize,
22    /// Number of parameters in each tensor
23    pub param_size: Vec<usize>,
24    /// Number of parameter tensors
25    pub num_params: usize,
26    /// Gradient magnitude multiplier for extreme conditions
27    pub gradient_scale: f32,
28    /// Whether to test with infinite/NaN gradients
29    pub test_edge_cases: bool,
30    /// Maximum allowed execution time
31    pub max_execution_time: Duration,
32    /// Memory usage tracking
33    pub track_memory: bool,
34}
35
36impl Default for StressTestConfig {
37    fn default() -> Self {
38        Self {
39            num_steps: 1000,
40            param_size: vec![100, 100],
41            num_params: 10,
42            gradient_scale: 1.0,
43            test_edge_cases: true,
44            max_execution_time: Duration::from_secs(30),
45            track_memory: true,
46        }
47    }
48}
49
50#[allow(dead_code)]
51/// Results from stress testing
52#[derive(Debug, Clone)]
53pub struct StressTestResult {
54    /// Whether the test passed without errors
55    pub passed: bool,
56    /// Total execution time
57    pub execution_time: Duration,
58    /// Average time per step
59    pub avg_step_time: Duration,
60    /// Memory usage statistics
61    pub memory_stats: MemoryStats,
62    /// Performance metrics
63    pub performance_metrics: HashMap<String, f32>,
64    /// Any errors encountered
65    pub errors: Vec<String>,
66}
67
68#[allow(dead_code)]
69/// Memory usage statistics
70#[derive(Debug, Clone)]
71pub struct MemoryStats {
72    /// Peak memory usage (estimated)
73    pub peak_memory_mb: f32,
74    /// Average memory usage
75    pub avg_memory_mb: f32,
76    /// Memory growth rate
77    pub memory_growth_rate: f32,
78}
79
80impl Default for MemoryStats {
81    fn default() -> Self {
82        Self {
83            peak_memory_mb: 0.0,
84            avg_memory_mb: 0.0,
85            memory_growth_rate: 0.0,
86        }
87    }
88}
89
90#[allow(dead_code)]
91/// Stress tester for optimizers
92pub struct OptimizerStressTester {
93    config: StressTestConfig,
94}
95
96#[allow(dead_code)]
97impl OptimizerStressTester {
98    /// Create a new stress tester with configuration
99    pub fn new(config: StressTestConfig) -> Self {
100        Self { config }
101    }
102
103    /// Create a stress tester with default configuration
104    pub fn default() -> Self {
105        Self::new(StressTestConfig::default())
106    }
107
108    /// Run comprehensive stress tests on an optimizer
109    pub fn run_stress_test<O>(&self, mut optimizer: O) -> OptimizerResult<StressTestResult>
110    where
111        O: crate::Optimizer,
112    {
113        let start_time = Instant::now();
114        let mut errors = Vec::new();
115        let mut step_times = Vec::new();
116        let mut memory_measurements = Vec::new();
117
118        // Create large parameter tensors for stress testing
119        let mut params = Vec::new();
120        for i in 0..self.config.num_params {
121            let param = Arc::new(RwLock::new(randn::<f32>(&self.config.param_size).map_err(
122                |e| {
123                    OptimizerError::InvalidParameter(format!("Failed to create param {}: {}", i, e))
124                },
125            )?));
126            params.push(param);
127        }
128
129        // Run optimization steps with timing
130        for step in 0..self.config.num_steps {
131            let step_start = Instant::now();
132
133            // Set gradients for all parameters
134            for (i, param) in params.iter().enumerate() {
135                let gradient = if self.config.test_edge_cases && step % 100 == 50 {
136                    // Inject extreme gradients occasionally
137                    self.create_extreme_gradient(&self.config.param_size, step)?
138                } else {
139                    randn::<f32>(&self.config.param_size)
140                        .map_err(|e| {
141                            OptimizerError::InvalidParameter(format!(
142                                "Failed to create gradient for param {}: {}",
143                                i, e
144                            ))
145                        })?
146                        .mul_scalar(self.config.gradient_scale)
147                        .map_err(|e| {
148                            OptimizerError::InvalidParameter(format!(
149                                "Failed to scale gradient: {}",
150                                e
151                            ))
152                        })?
153                };
154
155                param.write().set_grad(Some(gradient));
156            }
157
158            // Perform optimization step
159            match optimizer.step() {
160                Ok(_) => {}
161                Err(e) => {
162                    errors.push(format!("Step {}: {}", step, e));
163                    if errors.len() > 10 {
164                        break; // Stop after too many errors
165                    }
166                }
167            }
168
169            let step_duration = step_start.elapsed();
170            step_times.push(step_duration);
171
172            // Memory tracking (simplified estimation)
173            if self.config.track_memory && step % 10 == 0 {
174                let estimated_memory = self.estimate_memory_usage(&params);
175                memory_measurements.push(estimated_memory);
176            }
177
178            // Check for timeout
179            if start_time.elapsed() > self.config.max_execution_time {
180                errors.push("Test exceeded maximum execution time".to_string());
181                break;
182            }
183        }
184
185        let total_time = start_time.elapsed();
186        let avg_step_time = if !step_times.is_empty() {
187            step_times.iter().sum::<Duration>() / step_times.len() as u32
188        } else {
189            Duration::from_nanos(0)
190        };
191
192        // Calculate memory statistics
193        let memory_stats = if self.config.track_memory && !memory_measurements.is_empty() {
194            let peak_memory = memory_measurements
195                .iter()
196                .fold(0.0f32, |acc, x| acc.max(*x));
197            let avg_memory =
198                memory_measurements.iter().sum::<f32>() / memory_measurements.len() as f32;
199            let growth_rate = if memory_measurements.len() > 1 {
200                (memory_measurements[memory_measurements.len() - 1] - memory_measurements[0])
201                    / memory_measurements.len() as f32
202            } else {
203                0.0
204            };
205
206            MemoryStats {
207                peak_memory_mb: peak_memory,
208                avg_memory_mb: avg_memory,
209                memory_growth_rate: growth_rate,
210            }
211        } else {
212            MemoryStats::default()
213        };
214
215        // Calculate performance metrics
216        let mut performance_metrics = HashMap::new();
217        performance_metrics.insert(
218            "steps_per_second".to_string(),
219            self.config.num_steps as f32 / total_time.as_secs_f32(),
220        );
221        performance_metrics.insert(
222            "error_rate".to_string(),
223            errors.len() as f32 / self.config.num_steps as f32,
224        );
225        if !step_times.is_empty() {
226            performance_metrics.insert(
227                "avg_step_time_ms".to_string(),
228                avg_step_time.as_millis() as f32,
229            );
230            performance_metrics.insert(
231                "max_step_time_ms".to_string(),
232                step_times
233                    .iter()
234                    .max()
235                    .expect("step_times is non-empty")
236                    .as_millis() as f32,
237            );
238        }
239
240        let passed = errors.is_empty() && total_time <= self.config.max_execution_time;
241
242        Ok(StressTestResult {
243            passed,
244            execution_time: total_time,
245            avg_step_time,
246            memory_stats,
247            performance_metrics,
248            errors,
249        })
250    }
251
252    /// Test optimizer stability under extreme conditions
253    pub fn test_extreme_conditions<O>(&self, mut optimizer: O) -> OptimizerResult<StressTestResult>
254    where
255        O: crate::Optimizer,
256    {
257        let start_time = Instant::now();
258        let mut errors = Vec::new();
259
260        // Create a single parameter for testing
261        let param = Arc::new(RwLock::new(zeros(&[10, 10])?));
262
263        // Test cases: [magnitude, description]
264        let test_cases = vec![
265            (1e10, "Very large gradients"),
266            (1e-10, "Very small gradients"),
267            (0.0, "Zero gradients"),
268            (f32::INFINITY, "Infinite gradients"),
269            (f32::NAN, "NaN gradients"),
270        ];
271
272        let test_cases_len = test_cases.len();
273        for (magnitude, description) in test_cases {
274            // Create gradient with the test magnitude
275            let mut grad_data = vec![magnitude; 100];
276            if magnitude.is_nan() {
277                grad_data = vec![f32::NAN; 100];
278            }
279
280            let grad_tensor = Tensor::from_vec(grad_data, &[10, 10]).map_err(|e| {
281                OptimizerError::InvalidParameter(format!("Failed to create test gradient: {}", e))
282            })?;
283
284            param.write().set_grad(Some(grad_tensor));
285
286            // Test optimizer step
287            match optimizer.step() {
288                Ok(_) => {
289                    // Check if parameters are still valid
290                    let param_values = param.read().to_vec().map_err(|e| {
291                        OptimizerError::InvalidParameter(format!(
292                            "Failed to read parameter values: {}",
293                            e
294                        ))
295                    })?;
296
297                    let has_invalid = param_values.iter().any(|&x| x.is_nan() || x.is_infinite());
298                    if has_invalid {
299                        errors.push(format!("{}: Parameters became invalid", description));
300                    }
301                }
302                Err(e) => {
303                    // Some errors are expected for extreme inputs
304                    if !matches!(magnitude, val if val.is_infinite() || val.is_nan()) {
305                        errors.push(format!("{}: Unexpected error: {}", description, e));
306                    }
307                }
308            }
309        }
310
311        let total_time = start_time.elapsed();
312        let passed = errors.len() < test_cases_len / 2; // Allow some failures for extreme cases
313
314        let mut performance_metrics = HashMap::new();
315        performance_metrics.insert(
316            "extreme_case_success_rate".to_string(),
317            (test_cases_len - errors.len()) as f32 / test_cases_len as f32,
318        );
319
320        Ok(StressTestResult {
321            passed,
322            execution_time: total_time,
323            avg_step_time: total_time / test_cases_len as u32,
324            memory_stats: MemoryStats::default(),
325            performance_metrics,
326            errors,
327        })
328    }
329
330    /// Create extreme gradients for testing edge cases
331    fn create_extreme_gradient(&self, shape: &[usize], step: usize) -> OptimizerResult<Tensor> {
332        let total_elements: usize = shape.iter().product();
333
334        let gradient_data = match step % 4 {
335            0 => vec![1e6; total_elements],  // Very large
336            1 => vec![1e-6; total_elements], // Very small
337            2 => vec![0.0; total_elements],  // Zero
338            _ => {
339                // Alternating pattern
340                (0..total_elements)
341                    .map(|i| if i % 2 == 0 { 1e3 } else { -1e3 })
342                    .collect()
343            }
344        };
345
346        Tensor::from_vec(gradient_data, shape).map_err(|e| {
347            OptimizerError::InvalidParameter(format!("Failed to create extreme gradient: {}", e))
348        })
349    }
350
351    /// Estimate memory usage of parameters (simplified)
352    fn estimate_memory_usage(&self, params: &[Arc<RwLock<Tensor>>]) -> f32 {
353        let mut total_elements = 0;
354        for param in params {
355            if let Some(param_read) = param.try_read() {
356                let shape = param_read.shape();
357                total_elements += shape.dims().iter().product::<usize>();
358            }
359        }
360
361        // Estimate: 4 bytes per f32 element, convert to MB
362        (total_elements * 4) as f32 / (1024.0 * 1024.0)
363    }
364}
365
366#[cfg(test)]
367mod tests {
368    use super::*;
369    use crate::{adam::Adam, sgd::SGD};
370
371    #[test]
372    fn test_stress_tester_creation() -> OptimizerResult<()> {
373        let config = StressTestConfig::default();
374        let _tester = OptimizerStressTester::new(config);
375        Ok(())
376    }
377
378    #[test]
379    fn test_basic_stress_test() -> OptimizerResult<()> {
380        let mut config = StressTestConfig::default();
381        config.num_steps = 10; // Keep test fast
382        config.num_params = 2;
383        config.param_size = vec![5, 5];
384
385        let tester = OptimizerStressTester::new(config);
386        let param = Arc::new(RwLock::new(randn::<f32>(&[5, 5])?));
387        let optimizer = SGD::new(vec![param], 0.01, None, None, None, false);
388
389        let result = tester.run_stress_test(optimizer)?;
390
391        // Should complete without major issues
392        assert!(result.execution_time.as_secs() < 5);
393        assert!(result.performance_metrics.contains_key("steps_per_second"));
394        Ok(())
395    }
396
397    #[test]
398    fn test_extreme_conditions() -> OptimizerResult<()> {
399        let config = StressTestConfig::default();
400        let tester = OptimizerStressTester::new(config);
401        let param = Arc::new(RwLock::new(zeros(&[10, 10])?));
402        let optimizer = Adam::new(vec![param], Some(0.01), None, None, None, false);
403
404        let result = tester.test_extreme_conditions(optimizer)?;
405
406        // Should handle at least some extreme cases
407        assert!(result
408            .performance_metrics
409            .contains_key("extreme_case_success_rate"));
410        Ok(())
411    }
412
413    #[test]
414    fn test_memory_estimation() -> OptimizerResult<()> {
415        let tester = OptimizerStressTester::default();
416        let params = vec![
417            Arc::new(RwLock::new(zeros(&[100, 100])?)),
418            Arc::new(RwLock::new(zeros(&[50, 50])?)),
419        ];
420
421        let memory_usage = tester.estimate_memory_usage(&params);
422
423        // Should estimate reasonable memory usage (100*100 + 50*50 = 12500 elements * 4 bytes)
424        assert!(memory_usage > 0.04); // At least 0.04 MB
425        assert!(memory_usage < 1.0); // Less than 1 MB
426        Ok(())
427    }
428}