1use crate::{
11 adagrad::AdaGrad, adam::Adam, rmsprop::RMSprop, sgd::SGD, Optimizer, OptimizerError,
12 OptimizerResult,
13};
14use parking_lot::RwLock;
15use std::ops::{Add, Mul, Sub};
16use std::sync::Arc;
17use torsh_core::{
18 device::{CpuDevice, Device, DeviceType},
19 DType,
20};
21use torsh_tensor::{
22 creation::{eye, randn, tensor_scalar, zeros},
23 Tensor,
24};
25
26#[derive(Debug, Clone)]
28pub struct StabilityTestConfig {
29 pub num_steps: usize,
31 pub tolerance: f32,
33 pub max_param_magnitude: f32,
35 pub min_progress: f32,
37 pub device: Arc<CpuDevice>,
39}
40
41impl Default for StabilityTestConfig {
42 fn default() -> Self {
43 Self {
44 num_steps: 100,
45 tolerance: 1e-6,
46 max_param_magnitude: 1e10,
47 min_progress: 1e-8,
48 device: Arc::new(CpuDevice::new()),
49 }
50 }
51}
52
53#[derive(Debug)]
55pub struct StabilityTestResult {
56 pub passed: bool,
58 pub final_loss: f32,
60 pub max_param_magnitude: f32,
62 pub nan_count: usize,
64 pub error_message: Option<String>,
66}
67
68pub struct NumericalStabilityTests {
70 config: StabilityTestConfig,
71}
72
73impl NumericalStabilityTests {
74 pub fn new() -> Self {
76 Self {
77 config: StabilityTestConfig::default(),
78 }
79 }
80
81 pub fn with_config(config: StabilityTestConfig) -> Self {
83 Self { config }
84 }
85
86 pub fn test_extreme_gradients<O: Optimizer>(
88 &self,
89 mut optimizer: O,
90 ) -> OptimizerResult<StabilityTestResult> {
91 let mut params = randn::<f32>(&[10, 10])?;
93 let mut max_param_magnitude = 0.0f32;
94 let mut nan_count = 0;
95
96 for step in 0..self.config.num_steps {
97 let grad_scale = 10.0f32.powi(step as i32 / 20); let grads = randn::<f32>(&[10, 10])?.mul_scalar(grad_scale)?;
100
101 let grad_data = grads.to_vec()?;
103 let has_nan_or_inf = grad_data
104 .iter()
105 .any(|&x: &f32| x.is_nan() || x.is_infinite());
106 if has_nan_or_inf {
107 nan_count += 1;
108 continue;
109 }
110
111 params.set_grad(Some(grads));
113 optimizer.step()?;
114
115 let param_norm = params.norm()?.to_vec()?[0];
117 max_param_magnitude = max_param_magnitude.max(param_norm);
118
119 let param_data = params.to_vec()?;
121 let has_nan_or_inf = param_data
122 .iter()
123 .any(|&x: &f32| x.is_nan() || x.is_infinite());
124 if has_nan_or_inf {
125 return Ok(StabilityTestResult {
126 passed: false,
127 final_loss: f32::NAN,
128 max_param_magnitude,
129 nan_count,
130 error_message: Some(format!("Parameters became NaN/infinite at step {}", step)),
131 });
132 }
133
134 if param_norm > self.config.max_param_magnitude {
136 return Ok(StabilityTestResult {
137 passed: false,
138 final_loss: param_norm,
139 max_param_magnitude,
140 nan_count,
141 error_message: Some(format!(
142 "Parameters exploded to magnitude {} at step {}",
143 param_norm, step
144 )),
145 });
146 }
147 }
148
149 Ok(StabilityTestResult {
150 passed: true,
151 final_loss: params.norm()?.item()?,
152 max_param_magnitude,
153 nan_count,
154 error_message: None,
155 })
156 }
157
158 pub fn test_ill_conditioned_quadratic<O: Optimizer>(
160 &self,
161 mut optimizer: O,
162 ) -> OptimizerResult<StabilityTestResult> {
163 let device = self.config.device.clone();
165 let dim = 10;
166
167 let mut hessian_data = vec![0.0f32; dim * dim];
169 for i in 0..dim {
170 let eigenval = if i == 0 { 1000.0 } else { 0.001 };
171 hessian_data[i * dim + i] = eigenval;
172 }
173 let hessian = Tensor::from_data(hessian_data, vec![dim, dim], DeviceType::Cpu)?;
174
175 let mut params = randn::<f32>(&[dim])?;
176 let mut initial_loss = f32::INFINITY;
177 let mut max_param_magnitude = 0.0f32;
178 let mut nan_count = 0;
179
180 for step in 0..self.config.num_steps {
181 let grads = hessian.matmul(¶ms.unsqueeze(1)?)?.squeeze(1)?;
183
184 let grad_data = grads.to_vec()?;
186 let has_nan_or_inf = grad_data
187 .iter()
188 .any(|&x: &f32| x.is_nan() || x.is_infinite());
189 if has_nan_or_inf {
190 nan_count += 1;
191 continue;
192 }
193
194 let loss = params
196 .unsqueeze(0)?
197 .matmul(&grads.unsqueeze(1)?)?
198 .squeeze_all()?
199 .mul_scalar(0.5)?
200 .to_vec()?[0];
201
202 if step == 0 {
203 initial_loss = loss;
204 }
205
206 params.set_grad(Some(grads));
208 optimizer.step()?;
209
210 let param_norm = params.norm()?.to_vec()?[0];
212 max_param_magnitude = max_param_magnitude.max(param_norm);
213
214 let param_data = params.to_vec()?;
216 let has_nan_or_inf = param_data
217 .iter()
218 .any(|&x: &f32| x.is_nan() || x.is_infinite());
219 if has_nan_or_inf {
220 return Ok(StabilityTestResult {
221 passed: false,
222 final_loss: f32::NAN,
223 max_param_magnitude,
224 nan_count,
225 error_message: Some(format!("Parameters became NaN/infinite at step {}", step)),
226 });
227 }
228
229 if param_norm > self.config.max_param_magnitude {
231 return Ok(StabilityTestResult {
232 passed: false,
233 final_loss: loss,
234 max_param_magnitude,
235 nan_count,
236 error_message: Some(format!(
237 "Parameters exploded to magnitude {} at step {}",
238 param_norm, step
239 )),
240 });
241 }
242 }
243
244 let final_loss = {
245 let grads = hessian.matmul(¶ms.unsqueeze(1)?)?.squeeze(1)?;
246 params
247 .unsqueeze(0)?
248 .matmul(&grads.unsqueeze(1)?)?
249 .squeeze_all()?
250 .mul_scalar(0.5)?
251 .to_vec()?[0]
252 };
253
254 let progress = (initial_loss - final_loss) / initial_loss.max(1e-8);
256 if progress < self.config.min_progress {
257 return Ok(StabilityTestResult {
258 passed: false,
259 final_loss,
260 max_param_magnitude,
261 nan_count,
262 error_message: Some(format!("Insufficient progress: {:.2e}", progress)),
263 });
264 }
265
266 Ok(StabilityTestResult {
267 passed: true,
268 final_loss,
269 max_param_magnitude,
270 nan_count,
271 error_message: None,
272 })
273 }
274
275 pub fn test_noisy_gradients<O: Optimizer>(
277 &self,
278 mut optimizer: O,
279 ) -> OptimizerResult<StabilityTestResult> {
280 let device = self.config.device.clone();
281 let mut params = randn::<f32>(&[50])?;
282 let target = zeros(&[50])?;
283
284 let mut max_param_magnitude = 0.0f32;
285 let mut nan_count = 0;
286 let mut initial_loss = f32::INFINITY;
287
288 for step in 0..self.config.num_steps {
289 let clean_grads = params.sub(&target)?;
291
292 let noise_scale = 0.1; let noise = randn::<f32>(&[50])?.mul_scalar(noise_scale)?;
295 let noisy_grads = clean_grads.add(&noise)?;
296
297 let noisy_grad_data = noisy_grads.to_vec()?;
299 let has_nan_or_inf = noisy_grad_data
300 .iter()
301 .any(|&x: &f32| x.is_nan() || x.is_infinite());
302 if has_nan_or_inf {
303 nan_count += 1;
304 continue;
305 }
306
307 let loss = params.sub(&target)?.pow(2.0)?.mean(None, false)?.to_vec()?[0];
309
310 if step == 0 {
311 initial_loss = loss;
312 }
313
314 params.set_grad(Some(noisy_grads));
316 optimizer.step()?;
317
318 let param_norm = params.norm()?.to_vec()?[0];
320 max_param_magnitude = max_param_magnitude.max(param_norm);
321
322 let param_data = params.to_vec()?;
324 let has_nan_or_inf = param_data
325 .iter()
326 .any(|&x: &f32| x.is_nan() || x.is_infinite());
327 if has_nan_or_inf {
328 return Ok(StabilityTestResult {
329 passed: false,
330 final_loss: f32::NAN,
331 max_param_magnitude,
332 nan_count,
333 error_message: Some(format!("Parameters became NaN/infinite at step {}", step)),
334 });
335 }
336
337 if param_norm > self.config.max_param_magnitude {
339 return Ok(StabilityTestResult {
340 passed: false,
341 final_loss: loss,
342 max_param_magnitude,
343 nan_count,
344 error_message: Some(format!(
345 "Parameters exploded to magnitude {} at step {}",
346 param_norm, step
347 )),
348 });
349 }
350 }
351
352 let final_loss = params.sub(&target)?.pow(2.0)?.mean(None, false)?.item()?;
353
354 let progress = (initial_loss - final_loss) / initial_loss.max(1e-8);
356 if progress < self.config.min_progress {
357 return Ok(StabilityTestResult {
358 passed: false,
359 final_loss,
360 max_param_magnitude,
361 nan_count,
362 error_message: Some(format!(
363 "Insufficient progress with noisy gradients: {:.2e}",
364 progress
365 )),
366 });
367 }
368
369 Ok(StabilityTestResult {
370 passed: true,
371 final_loss,
372 max_param_magnitude,
373 nan_count,
374 error_message: None,
375 })
376 }
377
378 pub fn test_sparse_gradients<O: Optimizer>(
380 &self,
381 mut optimizer: O,
382 ) -> OptimizerResult<StabilityTestResult> {
383 let device = self.config.device.clone();
384 let mut params = randn::<f32>(&[100])?;
385
386 let mut max_param_magnitude = 0.0f32;
387 let mut nan_count = 0;
388
389 for step in 0..self.config.num_steps {
390 let mut grads_data = vec![0.0f32; 100];
392
393 let sparsity = 0.1; for i in 0..100 {
396 if (i * 17 + step) % 10 == 0 {
398 let grad_val = ((i as f32 * 0.1) % 2.0) - 1.0; grads_data[i] = grad_val;
401 }
402 }
403 let grads = Tensor::from_data(grads_data, vec![100], DeviceType::Cpu)?;
404
405 let grad_data = grads.to_vec()?;
407 let has_nan_or_inf = grad_data
408 .iter()
409 .any(|&x: &f32| x.is_nan() || x.is_infinite());
410 if has_nan_or_inf {
411 nan_count += 1;
412 continue;
413 }
414
415 params.set_grad(Some(grads));
417 optimizer.step()?;
418
419 let param_norm = params.norm()?.to_vec()?[0];
421 max_param_magnitude = max_param_magnitude.max(param_norm);
422
423 let param_data = params.to_vec()?;
425 let has_nan_or_inf = param_data
426 .iter()
427 .any(|&x: &f32| x.is_nan() || x.is_infinite());
428 if has_nan_or_inf {
429 return Ok(StabilityTestResult {
430 passed: false,
431 final_loss: f32::NAN,
432 max_param_magnitude,
433 nan_count,
434 error_message: Some(format!("Parameters became NaN/infinite at step {}", step)),
435 });
436 }
437
438 if param_norm > self.config.max_param_magnitude {
440 return Ok(StabilityTestResult {
441 passed: false,
442 final_loss: param_norm,
443 max_param_magnitude,
444 nan_count,
445 error_message: Some(format!(
446 "Parameters exploded to magnitude {} at step {}",
447 param_norm, step
448 )),
449 });
450 }
451 }
452
453 Ok(StabilityTestResult {
454 passed: true,
455 final_loss: params.norm()?.item()?,
456 max_param_magnitude,
457 nan_count,
458 error_message: None,
459 })
460 }
461
462 pub fn run_single_test<O: Optimizer>(
465 &self,
466 optimizer: O,
467 test_name: &str,
468 ) -> OptimizerResult<StabilityTestResult> {
469 match test_name {
470 "extreme_gradients" => self.test_extreme_gradients(optimizer),
471 "ill_conditioned_quadratic" => self.test_ill_conditioned_quadratic(optimizer),
472 "noisy_gradients" => self.test_noisy_gradients(optimizer),
473 "sparse_gradients" => self.test_sparse_gradients(optimizer),
474 _ => Err(OptimizerError::InvalidParameter(format!(
475 "Unknown test: {}",
476 test_name
477 ))),
478 }
479 }
480}
481
482pub fn run_comprehensive_stability_tests() -> OptimizerResult<()> {
484 let test_suite = NumericalStabilityTests::new();
485
486 let adam_params = randn::<f32>(&[10, 10])?;
488 let adam = Adam::new(
489 vec![Arc::new(RwLock::new(adam_params))],
490 Some(0.001),
491 None,
492 None,
493 None,
494 false,
495 );
496
497 println!("Testing Adam optimizer stability with extreme gradients...");
498 let adam_result = test_suite.run_single_test(adam, "extreme_gradients")?;
499 println!(
500 " extreme_gradients: {}",
501 if adam_result.passed { "PASS" } else { "FAIL" }
502 );
503 if let Some(error) = adam_result.error_message {
504 println!(" Error: {}", error);
505 }
506
507 let sgd_params = randn::<f32>(&[10, 10])?;
509 let sgd = SGD::new(
510 vec![Arc::new(RwLock::new(sgd_params))],
511 0.01,
512 None,
513 None,
514 None,
515 false,
516 );
517
518 println!("\nTesting SGD optimizer stability with noisy gradients...");
519 let sgd_result = test_suite.run_single_test(sgd, "noisy_gradients")?;
520 println!(
521 " noisy_gradients: {}",
522 if sgd_result.passed { "PASS" } else { "FAIL" }
523 );
524 if let Some(error) = sgd_result.error_message {
525 println!(" Error: {}", error);
526 }
527
528 let rmsprop_params = randn::<f32>(&[10, 10])?;
530 let rmsprop = RMSprop::new(
531 vec![Arc::new(RwLock::new(rmsprop_params))],
532 Some(0.01),
533 None,
534 None,
535 None,
536 None,
537 false,
538 );
539
540 println!("\nTesting RMSprop optimizer stability with sparse gradients...");
541 let rmsprop_result = test_suite.run_single_test(rmsprop, "sparse_gradients")?;
542 println!(
543 " sparse_gradients: {}",
544 if rmsprop_result.passed {
545 "PASS"
546 } else {
547 "FAIL"
548 }
549 );
550 if let Some(error) = rmsprop_result.error_message {
551 println!(" Error: {}", error);
552 }
553
554 Ok(())
555}
556
557#[cfg(test)]
558mod tests {
559 use super::*;
560
561 #[test]
562 fn test_stability_test_config() {
563 let config = StabilityTestConfig::default();
564 assert_eq!(config.num_steps, 100);
565 assert_eq!(config.tolerance, 1e-6);
566 assert_eq!(config.max_param_magnitude, 1e10);
567 assert_eq!(config.min_progress, 1e-8);
568 }
569
570 #[test]
571 fn test_stability_test_result() {
572 let result = StabilityTestResult {
573 passed: true,
574 final_loss: 0.5,
575 max_param_magnitude: 10.0,
576 nan_count: 0,
577 error_message: None,
578 };
579
580 assert!(result.passed);
581 assert_eq!(result.final_loss, 0.5);
582 assert_eq!(result.max_param_magnitude, 10.0);
583 assert_eq!(result.nan_count, 0);
584 assert!(result.error_message.is_none());
585 }
586
587 #[test]
588 fn test_numerical_stability_tests_creation() {
589 let test_suite = NumericalStabilityTests::new();
590 assert_eq!(test_suite.config.num_steps, 100);
591
592 let custom_config = StabilityTestConfig {
593 num_steps: 50,
594 tolerance: 1e-5,
595 max_param_magnitude: 1e8,
596 min_progress: 1e-7,
597 device: Arc::new(CpuDevice::new()),
598 };
599
600 let custom_test_suite = NumericalStabilityTests::with_config(custom_config);
601 assert_eq!(custom_test_suite.config.num_steps, 50);
602 assert_eq!(custom_test_suite.config.tolerance, 1e-5);
603 }
604}