Skip to main content

torsh_optim/
cross_framework_validation.rs

1//! Cross-framework validation tests for ToRSh optimizers
2//!
3//! This module provides functionality to validate ToRSh optimizer behavior
4//! against other deep learning frameworks to ensure compatibility and correctness.
5
6use crate::{Optimizer, OptimizerError, OptimizerResult};
7use parking_lot::RwLock;
8use std::collections::HashMap;
9use std::sync::Arc;
10use torsh_tensor::{
11    creation::{randn, zeros},
12    Tensor,
13};
14
15#[allow(dead_code)]
16/// Cross-framework validation configuration
17#[derive(Debug, Clone)]
18pub struct ValidationConfig {
19    /// Tolerance for numerical differences
20    pub tolerance: f32,
21    /// Number of optimization steps to compare
22    pub num_steps: usize,
23    /// Learning rate for comparison
24    pub learning_rate: f32,
25    /// Whether to enable verbose logging
26    pub verbose: bool,
27}
28
29impl Default for ValidationConfig {
30    fn default() -> Self {
31        Self {
32            tolerance: 1e-4,
33            num_steps: 10,
34            learning_rate: 0.01,
35            verbose: false,
36        }
37    }
38}
39
40#[allow(dead_code)]
41/// Results from cross-framework validation
42#[derive(Debug, Clone)]
43pub struct ValidationResult {
44    /// Whether the validation passed
45    pub passed: bool,
46    /// Maximum difference observed
47    pub max_difference: f32,
48    /// Average difference across all steps
49    pub avg_difference: f32,
50    /// Per-step differences
51    pub step_differences: Vec<f32>,
52    /// Additional metrics
53    pub metrics: HashMap<String, f32>,
54}
55
56#[allow(dead_code)]
57/// Cross-framework validator for optimizers
58pub struct CrossFrameworkValidator {
59    config: ValidationConfig,
60}
61
62#[allow(dead_code)]
63impl CrossFrameworkValidator {
64    /// Create a new validator with the given configuration
65    pub fn new(config: ValidationConfig) -> Self {
66        Self { config }
67    }
68
69    /// Create a validator with default configuration
70    pub fn default() -> Self {
71        Self::new(ValidationConfig::default())
72    }
73
74    /// Validate an optimizer against PyTorch's equivalent
75    pub fn validate_against_pytorch<O>(
76        &self,
77        mut torsh_optimizer: O,
78        pytorch_reference: &[f32],
79    ) -> OptimizerResult<ValidationResult>
80    where
81        O: crate::Optimizer,
82    {
83        let mut differences = Vec::new();
84        let mut max_diff = 0.0f32;
85        let mut sum_diff = 0.0f32;
86
87        // Create test parameters
88        let param = Arc::new(RwLock::new(zeros(&[2, 2])?));
89
90        for step in 0..self.config.num_steps {
91            // Set a consistent gradient for testing
92            let grad_data = vec![0.1, 0.2, 0.3, 0.4];
93            let grad_tensor = Tensor::from_vec(grad_data, &[2, 2])?;
94            param.write().set_grad(Some(grad_tensor));
95
96            // Step the ToRSh optimizer
97            torsh_optimizer.step()?;
98
99            // Get the current parameter values
100            let torsh_values = param.read().to_vec()?;
101
102            // Compare with PyTorch reference (simplified for example)
103            let pytorch_values = &pytorch_reference[step * 4..(step + 1) * 4];
104
105            // Compute differences
106            let step_diff = torsh_values
107                .iter()
108                .zip(pytorch_values.iter())
109                .map(|(a, b)| (a - b).abs())
110                .fold(0.0f32, |acc, x| acc.max(x));
111
112            differences.push(step_diff);
113            max_diff = max_diff.max(step_diff);
114            sum_diff += step_diff;
115
116            if self.config.verbose {
117                println!("Step {}: max_diff = {:.6}", step, step_diff);
118            }
119        }
120
121        let avg_diff = sum_diff / self.config.num_steps as f32;
122        let passed = max_diff < self.config.tolerance;
123
124        let mut metrics = HashMap::new();
125        metrics.insert("convergence_rate".to_string(), avg_diff);
126        metrics.insert("stability_score".to_string(), 1.0 / (1.0 + max_diff));
127
128        Ok(ValidationResult {
129            passed,
130            max_difference: max_diff,
131            avg_difference: avg_diff,
132            step_differences: differences,
133            metrics,
134        })
135    }
136
137    /// Validate optimizer convergence properties
138    pub fn validate_convergence<O>(&self, mut optimizer: O) -> OptimizerResult<ValidationResult>
139    where
140        O: crate::Optimizer,
141    {
142        let mut losses = Vec::new();
143
144        // Create a simple quadratic optimization problem: min 0.5 * x^T * A * x
145        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 1])?));
146
147        for step in 0..self.config.num_steps {
148            // Compute gradient: grad = A * x where A is identity for simplicity
149            let current = param.read().clone();
150            let loss = current.pow(2.0)?.sum()?.to_vec()?[0] * 0.5;
151            losses.push(loss);
152
153            // Set gradient for the optimizer
154            param.write().set_grad(Some(current.clone()));
155
156            // Optimization step
157            optimizer.step()?;
158
159            if self.config.verbose {
160                println!("Step {}: loss = {:.6}", step, loss);
161            }
162        }
163
164        // Check if loss is decreasing (convergence)
165        let initial_loss = losses[0];
166        let final_loss = losses[losses.len() - 1];
167        let loss_reduction = (initial_loss - final_loss) / initial_loss;
168
169        let passed = loss_reduction > 0.1; // At least 10% improvement
170
171        let mut metrics = HashMap::new();
172        metrics.insert("initial_loss".to_string(), initial_loss);
173        metrics.insert("final_loss".to_string(), final_loss);
174        metrics.insert("loss_reduction".to_string(), loss_reduction);
175
176        Ok(ValidationResult {
177            passed,
178            max_difference: final_loss,
179            avg_difference: losses.iter().sum::<f32>() / losses.len() as f32,
180            step_differences: losses,
181            metrics,
182        })
183    }
184
185    /// Run a comprehensive validation suite
186    pub fn run_validation_suite<O>(
187        &self,
188        optimizer: O,
189    ) -> OptimizerResult<HashMap<String, ValidationResult>>
190    where
191        O: crate::Optimizer,
192    {
193        let mut results = HashMap::new();
194
195        // Test 1: Convergence validation
196        let convergence_result = self.validate_convergence(optimizer)?;
197        results.insert("convergence".to_string(), convergence_result);
198
199        // Note: We can't run multiple tests with the same optimizer since it doesn't implement Clone
200        // This is a limitation of the current design - each test consumes the optimizer
201
202        Ok(results)
203    }
204
205    /// Validate basic gradient descent properties
206    fn validate_gradient_descent_properties<O>(
207        &self,
208        _optimizer: O,
209    ) -> OptimizerResult<ValidationResult>
210    where
211        O: crate::Optimizer,
212    {
213        // Test that parameters move in the opposite direction of gradients
214        let param = Arc::new(RwLock::new(zeros(&[2, 2])?));
215        let initial_params = param.read().to_vec()?;
216
217        // Create a new optimizer with our test parameter
218        let mut optimizer = crate::SGD::new(vec![param.clone()], 0.1, None, None, None, false);
219
220        // Set positive gradients
221        let grad_tensor = Tensor::from_vec(vec![1.0, 1.0, 1.0, 1.0], &[2, 2])?;
222        param.write().set_grad(Some(grad_tensor));
223
224        optimizer.step()?;
225
226        let final_params = param.read().to_vec()?;
227
228        // Parameters should have moved in negative direction (opposite to gradient)
229        let moved_correctly = initial_params
230            .iter()
231            .zip(final_params.iter())
232            .all(|(initial, final_val)| final_val < initial);
233
234        let max_movement = initial_params
235            .iter()
236            .zip(final_params.iter())
237            .map(|(initial, final_val)| ((*initial - *final_val) as f32).abs())
238            .fold(0.0f32, |acc, x| acc.max(x));
239
240        let mut metrics = HashMap::new();
241        metrics.insert("max_movement".to_string(), max_movement);
242        metrics.insert(
243            "correct_direction".to_string(),
244            if moved_correctly { 1.0 } else { 0.0 },
245        );
246
247        Ok(ValidationResult {
248            passed: moved_correctly,
249            max_difference: max_movement,
250            avg_difference: max_movement / 4.0, // 4 parameters
251            step_differences: vec![max_movement],
252            metrics,
253        })
254    }
255}
256
257#[cfg(test)]
258mod tests {
259    use super::*;
260    use crate::{adam::Adam, sgd::SGD};
261
262    #[test]
263    fn test_cross_framework_validator_creation() -> OptimizerResult<()> {
264        let config = ValidationConfig::default();
265        let _validator = CrossFrameworkValidator::new(config);
266        Ok(())
267    }
268
269    #[test]
270    fn test_convergence_validation() -> OptimizerResult<()> {
271        let validator = CrossFrameworkValidator::default();
272        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 1])?));
273        let optimizer = SGD::new(vec![param], 0.01, None, None, None, false);
274
275        let result = validator.validate_convergence(optimizer)?;
276
277        // Should show some form of optimization progress
278        assert!(result.metrics.contains_key("loss_reduction"));
279        Ok(())
280    }
281
282    #[test]
283    fn test_gradient_descent_properties() -> OptimizerResult<()> {
284        let validator = CrossFrameworkValidator::default();
285        let param = Arc::new(RwLock::new(zeros(&[2, 2])?));
286        let optimizer = SGD::new(vec![param], 0.1, None, None, None, false);
287
288        let result = validator.validate_gradient_descent_properties(optimizer)?;
289
290        // Should move parameters in correct direction
291        assert_eq!(result.metrics.get("correct_direction"), Some(&1.0));
292        Ok(())
293    }
294
295    #[test]
296    fn test_validation_suite() -> OptimizerResult<()> {
297        let validator = CrossFrameworkValidator::default();
298        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 1])?));
299        let optimizer = Adam::new(vec![param], None, None, None, None, false);
300
301        let results = validator.run_validation_suite(optimizer)?;
302
303        assert!(results.contains_key("convergence"));
304        // Note: Only convergence test is run now since optimizer doesn't implement Clone
305        Ok(())
306    }
307}