Skip to main content

torsh_optim/
numerical_stability_tests.rs

1//! Numerical stability tests for optimizers
2//!
3//! This module provides comprehensive tests to ensure optimizers maintain
4//! numerical stability under various conditions including:
5//! - Extreme gradients (very large/small)
6//! - Ill-conditioned optimization landscapes
7//! - Different precision levels
8//! - Pathological cases
9
10use 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/// Test configuration for numerical stability
27#[derive(Debug, Clone)]
28pub struct StabilityTestConfig {
29    /// Number of optimization steps to run
30    pub num_steps: usize,
31    /// Tolerance for checking stability
32    pub tolerance: f32,
33    /// Maximum allowed parameter magnitude
34    pub max_param_magnitude: f32,
35    /// Minimum required progress (to avoid stagnation)
36    pub min_progress: f32,
37    /// Device to run tests on
38    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/// Result of a numerical stability test
54#[derive(Debug)]
55pub struct StabilityTestResult {
56    /// Whether the test passed
57    pub passed: bool,
58    /// Final loss value
59    pub final_loss: f32,
60    /// Maximum parameter magnitude encountered
61    pub max_param_magnitude: f32,
62    /// Number of NaN/infinite values encountered
63    pub nan_count: usize,
64    /// Detailed error message if test failed
65    pub error_message: Option<String>,
66}
67
68/// Numerical stability test suite
69pub struct NumericalStabilityTests {
70    config: StabilityTestConfig,
71}
72
73impl NumericalStabilityTests {
74    /// Create a new test suite with default configuration
75    pub fn new() -> Self {
76        Self {
77            config: StabilityTestConfig::default(),
78        }
79    }
80
81    /// Create a new test suite with custom configuration
82    pub fn with_config(config: StabilityTestConfig) -> Self {
83        Self { config }
84    }
85
86    /// Test optimizer with extreme gradients
87    pub fn test_extreme_gradients<O: Optimizer>(
88        &self,
89        mut optimizer: O,
90    ) -> OptimizerResult<StabilityTestResult> {
91        // Create parameters with reasonable initial values
92        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            // Create extreme gradients that increase with steps
98            let grad_scale = 10.0f32.powi(step as i32 / 20); // Exponentially increasing
99            let grads = randn::<f32>(&[10, 10])?.mul_scalar(grad_scale)?;
100
101            // Check for NaN/infinite gradients
102            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            // Apply gradients
112            params.set_grad(Some(grads));
113            optimizer.step()?;
114
115            // Check parameter stability
116            let param_norm = params.norm()?.to_vec()?[0];
117            max_param_magnitude = max_param_magnitude.max(param_norm);
118
119            // Check for NaN/infinite parameters
120            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            // Check for parameter explosion
135            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    /// Test optimizer with ill-conditioned quadratic function
159    pub fn test_ill_conditioned_quadratic<O: Optimizer>(
160        &self,
161        mut optimizer: O,
162    ) -> OptimizerResult<StabilityTestResult> {
163        // Create an ill-conditioned quadratic: f(x) = 0.5 * x^T * A * x where A has poor condition number
164        let device = self.config.device.clone();
165        let dim = 10;
166
167        // Create a matrix with poor condition number
168        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            // Compute gradients: grad = A * x
182            let grads = hessian.matmul(&params.unsqueeze(1)?)?.squeeze(1)?;
183
184            // Check for NaN/infinite gradients
185            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            // Compute loss: 0.5 * x^T * A * x
195            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            // Apply gradients
207            params.set_grad(Some(grads));
208            optimizer.step()?;
209
210            // Check parameter stability
211            let param_norm = params.norm()?.to_vec()?[0];
212            max_param_magnitude = max_param_magnitude.max(param_norm);
213
214            // Check for NaN/infinite parameters
215            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            // Check for parameter explosion
230            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(&params.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        // Check if we made sufficient progress
255        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    /// Test optimizer with noisy gradients
276    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            // Compute clean gradients (towards target)
290            let clean_grads = params.sub(&target)?;
291
292            // Add high-frequency noise
293            let noise_scale = 0.1; // 10% noise
294            let noise = randn::<f32>(&[50])?.mul_scalar(noise_scale)?;
295            let noisy_grads = clean_grads.add(&noise)?;
296
297            // Check for NaN/infinite gradients
298            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            // Compute loss
308            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            // Apply gradients
315            params.set_grad(Some(noisy_grads));
316            optimizer.step()?;
317
318            // Check parameter stability
319            let param_norm = params.norm()?.to_vec()?[0];
320            max_param_magnitude = max_param_magnitude.max(param_norm);
321
322            // Check for NaN/infinite parameters
323            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            // Check for parameter explosion
338            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        // Check convergence despite noise
355        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    /// Test optimizer with sparse gradients
379    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            // Create sparse gradients (only update 10% of parameters each step)
391            let mut grads_data = vec![0.0f32; 100];
392
393            // Set gradients for a random subset of parameters
394            let sparsity = 0.1; // 10% non-zero gradients
395            for i in 0..100 {
396                // Use deterministic pattern for testing (can be replaced with proper RNG)
397                if (i * 17 + step) % 10 == 0 {
398                    // Deterministic "random" pattern
399                    let grad_val = ((i as f32 * 0.1) % 2.0) - 1.0; // Deterministic gradient in [-1, 1]
400                    grads_data[i] = grad_val;
401                }
402            }
403            let grads = Tensor::from_data(grads_data, vec![100], DeviceType::Cpu)?;
404
405            // Check for NaN/infinite gradients
406            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            // Apply gradients
416            params.set_grad(Some(grads));
417            optimizer.step()?;
418
419            // Check parameter stability
420            let param_norm = params.norm()?.to_vec()?[0];
421            max_param_magnitude = max_param_magnitude.max(param_norm);
422
423            // Check for NaN/infinite parameters
424            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            // Check for parameter explosion
439            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    /// Run all stability tests for a given optimizer
463    /// Note: This consumes the optimizer since optimizers don't implement Clone
464    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
482/// Comprehensive test suite for common optimizers
483pub fn run_comprehensive_stability_tests() -> OptimizerResult<()> {
484    let test_suite = NumericalStabilityTests::new();
485
486    // Test Adam optimizer with extreme gradients
487    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    // Test SGD optimizer with noisy gradients
508    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    // Test RMSprop optimizer with sparse gradients
529    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}