1use crate::error::{OptimError, Result};
7use crate::optimizers::Optimizer;
8use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
9use scirs2_core::numeric::Float;
10use std::fmt::Debug;
11use std::marker::PhantomData;
12
13pub struct SAM<A, O, D>
56where
57 A: Float + ScalarOperand + Debug,
58 O: Optimizer<A, D> + Clone,
59 D: Dimension,
60{
61 inner_optimizer: O,
63 rho: A,
65 epsilon: A,
67 adaptive: bool,
69 perturbed_params: Option<Array<A, D>>,
71 original_params: Option<Array<A, D>>,
73 _phantom: PhantomData<D>,
75}
76
77impl<A, O, D> SAM<A, O, D>
78where
79 A: Float + ScalarOperand + Debug,
80 O: Optimizer<A, D> + Clone,
81 D: Dimension,
82{
83 pub fn new(inner_optimizer: O) -> Self {
85 Self {
86 inner_optimizer,
87 rho: A::from(0.05).expect("SAM: default rho (0.05) must fit in A"),
88 epsilon: A::from(1e-12).expect("SAM: default epsilon (1e-12) must fit in A"),
89 adaptive: false,
90 perturbed_params: None,
91 original_params: None,
92 _phantom: PhantomData,
93 }
94 }
95
96 pub fn with_config(inner_optimizer: O, rho: A, adaptive: bool) -> Self {
98 Self {
99 inner_optimizer,
100 rho,
101 epsilon: A::from(1e-12).expect("SAM: default epsilon (1e-12) must fit in A"),
102 adaptive,
103 perturbed_params: None,
104 original_params: None,
105 _phantom: PhantomData,
106 }
107 }
108
109 pub fn with_rho(mut self, rho: A) -> Self {
111 self.rho = rho;
112 self
113 }
114
115 pub fn with_epsilon(mut self, epsilon: A) -> Self {
117 self.epsilon = epsilon;
118 self
119 }
120
121 pub fn with_adaptive(mut self, adaptive: bool) -> Self {
123 self.adaptive = adaptive;
124 self
125 }
126
127 pub fn inner_optimizer(&self) -> &O {
129 &self.inner_optimizer
130 }
131
132 pub fn inner_optimizer_mut(&mut self) -> &mut O {
134 &mut self.inner_optimizer
135 }
136
137 pub fn rho(&self) -> A {
139 self.rho
140 }
141
142 pub fn epsilon(&self) -> A {
144 self.epsilon
145 }
146
147 pub fn is_adaptive(&self) -> bool {
149 self.adaptive
150 }
151
152 pub fn first_step(
163 &mut self,
164 params: &Array<A, D>,
165 gradients: &Array<A, D>,
166 ) -> Result<(Array<A, D>, A)> {
167 self.original_params = Some(params.clone());
169
170 let grad_norm = calculate_norm(gradients)?;
172
173 if grad_norm.is_zero() || !grad_norm.is_finite() {
174 return Err(OptimError::OptimizationError(
175 "Gradient norm is zero or not finite".to_string(),
176 ));
177 }
178
179 let e_w = if self.adaptive {
181 let param_norm = calculate_norm(params)?;
184 if param_norm.is_zero() || !param_norm.is_finite() {
185 let perturb = gradients / (grad_norm + self.epsilon);
187 &perturb * self.rho
188 } else {
189 let mut perturb = params.mapv(|p| p.abs() + self.epsilon);
191 perturb = &perturb / param_norm; gradients * &perturb * self.rho
194 }
195 } else {
196 let perturb = gradients / (grad_norm + self.epsilon);
198 &perturb * self.rho
199 };
200
201 let perturbed_params = params + &e_w;
203 self.perturbed_params = Some(perturbed_params.clone());
204
205 Ok((perturbed_params, calculate_norm(&e_w)?))
207 }
208
209 pub fn second_step(
220 &mut self,
221 params: &Array<A, D>,
222 gradients: &Array<A, D>,
223 ) -> Result<Array<A, D>> {
224 let original_params = match &self.original_params {
226 Some(_params) => params,
227 None => {
228 return Err(OptimError::OptimizationError(
229 "Must call first_step before second_step".to_string(),
230 ))
231 }
232 };
233
234 let updated_params = self.inner_optimizer.step(original_params, gradients)?;
236
237 self.perturbed_params = None;
239 self.original_params = None;
240
241 Ok(updated_params)
242 }
243
244 pub fn reset(&mut self) {
246 self.perturbed_params = None;
247 self.original_params = None;
248 }
249}
250
251impl<A, O, D> Clone for SAM<A, O, D>
252where
253 A: Float + ScalarOperand + Debug,
254 O: Optimizer<A, D> + Clone,
255 D: Dimension,
256{
257 fn clone(&self) -> Self {
258 Self {
259 inner_optimizer: self.inner_optimizer.clone(),
260 rho: self.rho,
261 epsilon: self.epsilon,
262 adaptive: self.adaptive,
263 perturbed_params: self.perturbed_params.clone(),
264 original_params: self.original_params.clone(),
265 _phantom: PhantomData,
266 }
267 }
268}
269
270impl<A, O, D> Debug for SAM<A, O, D>
271where
272 A: Float + ScalarOperand + Debug,
273 O: Optimizer<A, D> + Clone + Debug,
274 D: Dimension,
275{
276 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
277 f.debug_struct("SAM")
278 .field("inner_optimizer", &self.inner_optimizer)
279 .field("rho", &self.rho)
280 .field("epsilon", &self.epsilon)
281 .field("adaptive", &self.adaptive)
282 .finish()
283 }
284}
285
286impl<A, O, D> Optimizer<A, D> for SAM<A, O, D>
287where
288 A: Float + ScalarOperand + Debug + Send + Sync,
289 O: Optimizer<A, D> + Clone + Send + Sync,
290 D: Dimension,
291{
292 fn step(&mut self, _params: &Array<A, D>, _gradients: &Array<A, D>) -> Result<Array<A, D>> {
312 Err(OptimError::InvalidConfig(
313 "SAM requires two gradient evaluations and does not support the single-step \
314 Optimizer::step API. Call first_step(params, grads) to get the perturbed \
315 parameters, recompute the gradients there, then call \
316 second_step(params, perturbed_grads)."
317 .to_string(),
318 ))
319 }
320
321 fn set_learning_rate(&mut self, learning_rate: A) {
322 self.inner_optimizer.set_learning_rate(learning_rate);
323 }
324
325 fn get_learning_rate(&self) -> A {
326 self.inner_optimizer.get_learning_rate()
327 }
328}
329
330fn calculate_norm<A, D>(array: &Array<A, D>) -> Result<A>
332where
333 A: Float + ScalarOperand + Debug,
334 D: Dimension,
335{
336 let squared_sum = array.iter().fold(A::zero(), |acc, &x| acc + x * x);
337 let norm = squared_sum.sqrt();
338
339 if !norm.is_finite() {
340 return Err(OptimError::OptimizationError(
341 "Norm calculation resulted in non-finite value".to_string(),
342 ));
343 }
344
345 Ok(norm)
346}
347
348#[cfg(test)]
349mod tests {
350 use super::*;
351 use crate::optimizers::sgd::SGD;
352 use approx::assert_abs_diff_eq;
353 use scirs2_core::ndarray::Array1;
354
355 #[test]
356 fn test_sam_creation() {
357 let sgd = SGD::new(0.01);
358 let optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> = SAM::new(sgd);
359
360 assert_abs_diff_eq!(optimizer.rho(), 0.05);
361 assert_abs_diff_eq!(optimizer.get_learning_rate(), 0.01);
362 assert!(!optimizer.is_adaptive());
363 }
364
365 #[test]
366 fn test_sam_with_config() {
367 let sgd = SGD::new(0.01);
368 let optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> =
369 SAM::with_config(sgd, 0.1, true);
370
371 assert_abs_diff_eq!(optimizer.rho(), 0.1);
372 assert!(optimizer.is_adaptive());
373 }
374
375 #[test]
376 fn test_sam_first_step() {
377 let sgd = SGD::new(0.1);
378 let mut optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> = SAM::new(sgd);
379
380 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
381 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
382
383 let grad_norm = (0.1f64.powi(2) + 0.2f64.powi(2) + 0.3f64.powi(2)).sqrt();
385 let normalized_grads = gradients.mapv(|g| g / grad_norm);
386 let expected_perturb = normalized_grads.mapv(|g| g * 0.05);
387 let expected_params = ¶ms + &expected_perturb;
388
389 let (perturbed_params, perturb_size) = optimizer
390 .first_step(¶ms, &gradients)
391 .expect("first_step succeeds in test_sam_first_step");
392
393 assert_abs_diff_eq!(perturbed_params[0], expected_params[0], epsilon = 1e-6);
395 assert_abs_diff_eq!(perturbed_params[1], expected_params[1], epsilon = 1e-6);
396 assert_abs_diff_eq!(perturbed_params[2], expected_params[2], epsilon = 1e-6);
397
398 assert_abs_diff_eq!(perturb_size, 0.05, epsilon = 1e-6);
400 }
401
402 #[test]
403 fn test_sam_adaptive() {
404 let sgd = SGD::new(0.1);
405 let mut optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> =
406 SAM::with_config(sgd, 0.05, true);
407
408 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
409 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
410
411 let (perturbed_params, perturb_size) = optimizer
413 .first_step(¶ms, &gradients)
414 .expect("first_step succeeds in test_sam_adaptive");
415
416 assert!(perturb_size > 0.0 && perturb_size < 1.0); assert!(perturbed_params[0] != params[0]);
421 assert!(perturbed_params[1] != params[1]);
422 assert!(perturbed_params[2] != params[2]);
423
424 let delta0 = (perturbed_params[0] - params[0]).abs();
426 let delta2 = (perturbed_params[2] - params[2]).abs();
427 assert!(delta2 > delta0);
428 }
429
430 #[test]
431 fn test_sam_second_step() {
432 let sgd = SGD::new(0.1);
433 let mut optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> = SAM::new(sgd);
434
435 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
436 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
437
438 let _ = optimizer
440 .first_step(¶ms, &gradients)
441 .expect("first_step succeeds in test_sam_second_step");
442
443 let new_gradients = Array1::from_vec(vec![0.15, 0.25, 0.35]);
445
446 let updated_params = optimizer
448 .second_step(¶ms, &new_gradients)
449 .expect("second_step succeeds in test_sam_second_step");
450
451 let expected_params =
453 Array1::from_vec(vec![1.0 - 0.1 * 0.15, 2.0 - 0.1 * 0.25, 3.0 - 0.1 * 0.35]);
454
455 assert_abs_diff_eq!(updated_params[0], expected_params[0], epsilon = 1e-6);
456 assert_abs_diff_eq!(updated_params[1], expected_params[1], epsilon = 1e-6);
457 assert_abs_diff_eq!(updated_params[2], expected_params[2], epsilon = 1e-6);
458 }
459
460 #[test]
461 fn test_sam_reset() {
462 let sgd = SGD::new(0.1);
463 let mut optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> = SAM::new(sgd);
464
465 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
466 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
467
468 let _ = optimizer
470 .first_step(¶ms, &gradients)
471 .expect("first_step succeeds in test_sam_reset");
472
473 optimizer.reset();
475
476 let result = optimizer.second_step(¶ms, &gradients);
478 assert!(result.is_err());
479 }
480
481 #[test]
482 fn test_sam_error_handling() {
483 let sgd = SGD::new(0.1);
484 let mut optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> = SAM::new(sgd);
485
486 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
488 let zero_gradients = Array1::zeros(3);
489
490 let result = optimizer.first_step(¶ms, &zero_gradients);
491 assert!(result.is_err());
492 }
493
494 #[test]
500 fn test_sam_optimizer_step_is_rejected() {
501 let sgd = SGD::new(0.1);
502 let mut optimizer: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> = SAM::new(sgd);
503
504 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
505 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
506
507 let result =
508 Optimizer::<f64, scirs2_core::ndarray::Ix1>::step(&mut optimizer, ¶ms, &gradients);
509
510 let error = match result {
511 Ok(_) => panic!("SAM::step must not silently degenerate to the inner optimizer"),
512 Err(e) => e.to_string(),
513 };
514 assert!(
515 error.contains("first_step") && error.contains("second_step"),
516 "error message must point at the two-phase API, got: {error}"
517 );
518 }
519
520 #[test]
522 fn test_sam_two_phase_differs_from_inner_optimizer() {
523 let mut sam: SAM<f64, SGD<f64>, scirs2_core::ndarray::Ix1> = SAM::new(SGD::new(0.1));
524 let mut plain = SGD::new(0.1);
525
526 let params = Array1::from_vec(vec![1.0, 2.0, 3.0]);
527 let grads_at_theta = Array1::from_vec(vec![0.1, 0.2, 0.3]);
528 let grads_at_perturbed = Array1::from_vec(vec![0.15, 0.25, 0.35]);
530
531 let (_perturbed, _size) = sam
532 .first_step(¶ms, &grads_at_theta)
533 .expect("first_step failed");
534 let sam_out = sam
535 .second_step(¶ms, &grads_at_perturbed)
536 .expect("second_step failed");
537
538 let plain_out = plain
539 .step(¶ms, &grads_at_theta)
540 .expect("plain step failed");
541
542 assert!((sam_out[0] - plain_out[0]).abs() > 1e-9);
543 }
544}