1use 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#[derive(Debug, Clone)]
18pub struct ValidationConfig {
19 pub tolerance: f32,
21 pub num_steps: usize,
23 pub learning_rate: f32,
25 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#[derive(Debug, Clone)]
43pub struct ValidationResult {
44 pub passed: bool,
46 pub max_difference: f32,
48 pub avg_difference: f32,
50 pub step_differences: Vec<f32>,
52 pub metrics: HashMap<String, f32>,
54}
55
56#[allow(dead_code)]
57pub struct CrossFrameworkValidator {
59 config: ValidationConfig,
60}
61
62#[allow(dead_code)]
63impl CrossFrameworkValidator {
64 pub fn new(config: ValidationConfig) -> Self {
66 Self { config }
67 }
68
69 pub fn default() -> Self {
71 Self::new(ValidationConfig::default())
72 }
73
74 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 let param = Arc::new(RwLock::new(zeros(&[2, 2])?));
89
90 for step in 0..self.config.num_steps {
91 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 torsh_optimizer.step()?;
98
99 let torsh_values = param.read().to_vec()?;
101
102 let pytorch_values = &pytorch_reference[step * 4..(step + 1) * 4];
104
105 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 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 let param = Arc::new(RwLock::new(randn::<f32>(&[2, 1])?));
146
147 for step in 0..self.config.num_steps {
148 let current = param.read().clone();
150 let loss = current.pow(2.0)?.sum()?.to_vec()?[0] * 0.5;
151 losses.push(loss);
152
153 param.write().set_grad(Some(current.clone()));
155
156 optimizer.step()?;
158
159 if self.config.verbose {
160 println!("Step {}: loss = {:.6}", step, loss);
161 }
162 }
163
164 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; 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 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 let convergence_result = self.validate_convergence(optimizer)?;
197 results.insert("convergence".to_string(), convergence_result);
198
199 Ok(results)
203 }
204
205 fn validate_gradient_descent_properties<O>(
207 &self,
208 _optimizer: O,
209 ) -> OptimizerResult<ValidationResult>
210 where
211 O: crate::Optimizer,
212 {
213 let param = Arc::new(RwLock::new(zeros(&[2, 2])?));
215 let initial_params = param.read().to_vec()?;
216
217 let mut optimizer = crate::SGD::new(vec![param.clone()], 0.1, None, None, None, false);
219
220 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 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, 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 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 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 Ok(())
306 }
307}