1use 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#[derive(Debug, Clone)]
19pub struct StressTestConfig {
20 pub num_steps: usize,
22 pub param_size: Vec<usize>,
24 pub num_params: usize,
26 pub gradient_scale: f32,
28 pub test_edge_cases: bool,
30 pub max_execution_time: Duration,
32 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#[derive(Debug, Clone)]
53pub struct StressTestResult {
54 pub passed: bool,
56 pub execution_time: Duration,
58 pub avg_step_time: Duration,
60 pub memory_stats: MemoryStats,
62 pub performance_metrics: HashMap<String, f32>,
64 pub errors: Vec<String>,
66}
67
68#[allow(dead_code)]
69#[derive(Debug, Clone)]
71pub struct MemoryStats {
72 pub peak_memory_mb: f32,
74 pub avg_memory_mb: f32,
76 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)]
91pub struct OptimizerStressTester {
93 config: StressTestConfig,
94}
95
96#[allow(dead_code)]
97impl OptimizerStressTester {
98 pub fn new(config: StressTestConfig) -> Self {
100 Self { config }
101 }
102
103 pub fn default() -> Self {
105 Self::new(StressTestConfig::default())
106 }
107
108 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 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 for step in 0..self.config.num_steps {
131 let step_start = Instant::now();
132
133 for (i, param) in params.iter().enumerate() {
135 let gradient = if self.config.test_edge_cases && step % 100 == 50 {
136 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 match optimizer.step() {
160 Ok(_) => {}
161 Err(e) => {
162 errors.push(format!("Step {}: {}", step, e));
163 if errors.len() > 10 {
164 break; }
166 }
167 }
168
169 let step_duration = step_start.elapsed();
170 step_times.push(step_duration);
171
172 if self.config.track_memory && step % 10 == 0 {
174 let estimated_memory = self.estimate_memory_usage(¶ms);
175 memory_measurements.push(estimated_memory);
176 }
177
178 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 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 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 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 let param = Arc::new(RwLock::new(zeros(&[10, 10])?));
262
263 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 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 match optimizer.step() {
288 Ok(_) => {
289 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 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; 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 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], 1 => vec![1e-6; total_elements], 2 => vec![0.0; total_elements], _ => {
339 (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 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 (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; 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 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 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(¶ms);
422
423 assert!(memory_usage > 0.04); assert!(memory_usage < 1.0); Ok(())
427 }
428}