1use scirs2_core::ndarray::ScalarOperand;
7use scirs2_core::numeric::{Float, NumCast};
8use scirs2_core::random::{rngs::StdRng, seeded_rng, thread_rng, CoreRandom};
9use std::fmt::Debug;
10
11use super::LearningRateScheduler;
12
13type NoiseRng = CoreRandom<StdRng>;
15
16fn from_f64<A: Float>(v: f64) -> A {
18 <A as NumCast>::from(v).unwrap_or_else(A::zero)
19}
20
21fn from_usize<A: Float>(v: usize) -> A {
23 <A as NumCast>::from(v).unwrap_or_else(A::zero)
24}
25
26fn denom_from_usize<A: Float>(v: usize) -> A {
30 match <A as NumCast>::from(v) {
31 Some(x) if x != A::zero() => x,
32 _ => A::one(),
33 }
34}
35
36fn standard_normal(rng: &mut NoiseRng) -> f64 {
42 let u1: f64 = 1.0 - rng.gen_range(0.0f64..1.0f64);
43 let u2: f64 = rng.gen_range(0.0f64..1.0f64);
44 (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
45}
46
47#[derive(Debug, Clone, Copy)]
49pub enum NoiseDistribution<A: Float> {
50 Uniform {
52 min: A,
54 max: A,
56 },
57 Gaussian {
59 mean: A,
61 std_dev: A,
63 },
64 Cyclical {
66 amplitude: A,
68 period: usize,
70 },
71 Decaying {
73 initial_scale: A,
75 final_scale: A,
77 decay_steps: usize,
79 },
80}
81
82pub struct NoiseInjectionScheduler<A, S>
92where
93 A: Float + Debug + ScalarOperand,
94 S: LearningRateScheduler<A>,
95{
96 base_scheduler: S,
98 noise_dist: NoiseDistribution<A>,
100 step_count: usize,
102 seed: u64,
104 rng: NoiseRng,
106 current_noise: A,
108 min_lr: A,
110}
111
112impl<A, S> NoiseInjectionScheduler<A, S>
113where
114 A: Float + Debug + ScalarOperand,
115 S: LearningRateScheduler<A>,
116{
117 pub fn new(base_scheduler: S, noise_dist: NoiseDistribution<A>, min_lr: A) -> Self {
151 let seed: u64 = thread_rng().gen_range(0u64..u64::MAX);
152 Self::new_seeded(base_scheduler, noise_dist, min_lr, seed)
153 }
154
155 pub fn new_seeded(
175 base_scheduler: S,
176 noise_dist: NoiseDistribution<A>,
177 min_lr: A,
178 seed: u64,
179 ) -> Self {
180 let mut scheduler = Self {
181 base_scheduler,
182 noise_dist,
183 step_count: 0,
184 seed,
185 rng: seeded_rng(seed),
186 current_noise: A::zero(),
187 min_lr,
188 };
189 scheduler.current_noise = scheduler.sample_noise();
190 scheduler
191 }
192
193 pub fn with_seed(mut self, seed: u64) -> Self {
198 self.seed = seed;
199 self.rng = seeded_rng(seed);
200 self.current_noise = self.sample_noise();
201 self
202 }
203
204 pub fn seed(&self) -> u64 {
206 self.seed
207 }
208
209 pub fn current_noise(&self) -> A {
211 self.current_noise
212 }
213
214 fn sample_noise(&mut self) -> A {
216 match self.noise_dist {
217 NoiseDistribution::Uniform { min, max } => {
218 let min_f64 = min.to_f64().unwrap_or(0.0);
219 let max_f64 = max.to_f64().unwrap_or(0.0);
220 if !min_f64.is_finite() || !max_f64.is_finite() {
221 return A::zero();
222 }
223 if min_f64 >= max_f64 {
224 return from_f64::<A>(min_f64);
226 }
227 from_f64::<A>(self.rng.gen_range(min_f64..max_f64))
228 }
229 NoiseDistribution::Gaussian { mean, std_dev } => {
230 let mean_f64 = mean.to_f64().unwrap_or(0.0);
231 let std_dev_f64 = std_dev.to_f64().unwrap_or(0.0);
232 let z0 = standard_normal(&mut self.rng);
233 from_f64::<A>(mean_f64 + std_dev_f64 * z0)
234 }
235 NoiseDistribution::Cyclical { amplitude, period } => {
236 let period_f = denom_from_usize::<A>(period.max(1));
237 let step = from_usize::<A>(self.step_count);
238 let angle =
239 from_f64::<A>(2.0) * from_f64::<A>(std::f64::consts::PI) * (step / period_f);
240 amplitude * angle.sin()
241 }
242 NoiseDistribution::Decaying {
243 initial_scale,
244 final_scale,
245 decay_steps,
246 } => {
247 let decay_steps = decay_steps.max(1);
248 let decay_steps_a = denom_from_usize::<A>(decay_steps);
249 let step = from_usize::<A>(self.step_count.min(decay_steps));
250 let scale = initial_scale - (step / decay_steps_a) * (initial_scale - final_scale);
251
252 scale * from_f64::<A>(self.rng.gen_range(-1.0f64..1.0f64))
254 }
255 }
256 }
257}
258
259impl<A, S> LearningRateScheduler<A> for NoiseInjectionScheduler<A, S>
260where
261 A: Float + Debug + ScalarOperand,
262 S: LearningRateScheduler<A>,
263{
264 fn get_learning_rate(&self) -> A {
265 let base_lr = self.base_scheduler.get_learning_rate();
267 (base_lr + self.current_noise).max(self.min_lr)
268 }
269
270 fn step(&mut self) -> A {
271 self.base_scheduler.step();
273
274 self.step_count = self.step_count.saturating_add(1);
276 self.current_noise = self.sample_noise();
277
278 self.get_learning_rate()
279 }
280
281 fn reset(&mut self) {
282 self.base_scheduler.reset();
283 self.step_count = 0;
284 self.rng = seeded_rng(self.seed);
285 self.current_noise = self.sample_noise();
286 }
287}
288
289impl<A, S> Clone for NoiseInjectionScheduler<A, S>
291where
292 A: Float + Debug + ScalarOperand,
293 S: LearningRateScheduler<A> + Clone,
294{
295 fn clone(&self) -> Self {
296 Self {
297 base_scheduler: self.base_scheduler.clone(),
298 noise_dist: self.noise_dist,
299 step_count: self.step_count,
300 seed: self.seed,
301 rng: seeded_rng(self.seed),
305 current_noise: self.current_noise,
306 min_lr: self.min_lr,
307 }
308 }
309}
310
311impl<A, S> Debug for NoiseInjectionScheduler<A, S>
312where
313 A: Float + Debug + ScalarOperand,
314 S: LearningRateScheduler<A> + Debug,
315{
316 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
317 f.debug_struct("NoiseInjectionScheduler")
318 .field("base_scheduler", &self.base_scheduler)
319 .field("noise_dist", &self.noise_dist)
320 .field("step_count", &self.step_count)
321 .field("seed", &self.seed)
322 .field("current_noise", &self.current_noise)
323 .field("min_lr", &self.min_lr)
324 .finish()
325 }
326}
327
328#[cfg(test)]
329mod tests {
330 use super::*;
331 use crate::schedulers::ConstantScheduler;
332
333 #[test]
334 fn test_uniform_noise() {
335 let base_scheduler = ConstantScheduler::new(0.1);
337
338 let mut scheduler = NoiseInjectionScheduler::new(
340 base_scheduler,
341 NoiseDistribution::Uniform {
342 min: -0.02,
343 max: 0.02,
344 },
345 0.001,
346 );
347
348 let mut rates = Vec::with_capacity(100);
350 for _ in 0..100 {
351 rates.push(scheduler.step());
352 }
353
354 for &rate in &rates {
356 assert!((0.08..=0.12).contains(&rate));
357 }
358
359 let mean = rates.iter().sum::<f64>() / rates.len() as f64;
361 let variance = rates.iter().map(|&r| (r - mean).powi(2)).sum::<f64>() / rates.len() as f64;
362
363 assert!(variance > 0.0);
365 }
366
367 #[test]
368 fn test_gaussian_noise() {
369 let base_scheduler = ConstantScheduler::new(0.1);
370 let mut scheduler = NoiseInjectionScheduler::new(
371 base_scheduler,
372 NoiseDistribution::Gaussian {
373 mean: 0.0,
374 std_dev: 0.01,
375 },
376 0.001,
377 );
378
379 let mut rates = Vec::with_capacity(1000);
381 for _ in 0..1000 {
382 rates.push(scheduler.step());
383 }
384
385 assert!(rates.iter().all(|r| r.is_finite()));
387
388 let mean = rates.iter().sum::<f64>() / rates.len() as f64;
390
391 assert!((mean - 0.1).abs() < 0.01);
393 }
394
395 #[test]
396 fn test_cyclical_noise() {
397 let base_scheduler = ConstantScheduler::new(0.1);
398 let mut scheduler = NoiseInjectionScheduler::new(
399 base_scheduler,
400 NoiseDistribution::Cyclical {
401 amplitude: 0.05,
402 period: 10,
403 },
404 0.001,
405 );
406
407 let mut rates = Vec::with_capacity(20);
409 for _ in 0..20 {
410 rates.push(scheduler.step());
411 }
412
413 for i in 0..10 {
415 assert!((rates[i] - rates[i + 10]).abs() < 1e-10);
418 }
419 }
420
421 #[test]
422 fn test_decaying_noise() {
423 let base_scheduler = ConstantScheduler::new(0.1);
424 let mut scheduler = NoiseInjectionScheduler::new(
425 base_scheduler,
426 NoiseDistribution::Decaying {
427 initial_scale: 0.05,
428 final_scale: 0.001,
429 decay_steps: 100,
430 },
431 0.001,
432 );
433
434 let mut early_rates = Vec::with_capacity(50);
438 for _ in 0..50 {
439 early_rates.push(scheduler.step());
440 }
441 let early_mean = early_rates.iter().sum::<f64>() / early_rates.len() as f64;
442 let early_variance = early_rates
443 .iter()
444 .map(|&r| (r - early_mean).powi(2))
445 .sum::<f64>()
446 / early_rates.len() as f64;
447
448 let mut late_rates = Vec::with_capacity(50);
450 for _ in 0..50 {
451 late_rates.push(scheduler.step());
452 }
453 let late_mean = late_rates.iter().sum::<f64>() / late_rates.len() as f64;
454 let late_variance = late_rates
455 .iter()
456 .map(|&r| (r - late_mean).powi(2))
457 .sum::<f64>()
458 / late_rates.len() as f64;
459
460 assert!(early_variance > late_variance);
462 }
463
464 #[test]
465 fn test_min_lr() {
466 let base_scheduler = ConstantScheduler::new(0.01);
467 let mut scheduler = NoiseInjectionScheduler::new(
468 base_scheduler,
469 NoiseDistribution::Uniform {
470 min: -0.1, max: 0.0,
472 },
473 0.005, );
475
476 for _ in 0..100 {
478 assert!(scheduler.step() >= 0.005);
479 }
480 }
481
482 #[test]
483 fn test_get_learning_rate_is_idempotent() {
484 let mut scheduler = NoiseInjectionScheduler::new(
485 ConstantScheduler::new(0.1),
486 NoiseDistribution::Uniform {
487 min: -0.02,
488 max: 0.02,
489 },
490 0.001,
491 );
492
493 assert_eq!(scheduler.get_learning_rate(), scheduler.get_learning_rate());
494 for _ in 0..32 {
495 let stepped = scheduler.step();
496 assert_eq!(stepped, scheduler.get_learning_rate());
497 assert_eq!(stepped, scheduler.get_learning_rate());
498 }
499 }
500
501 #[test]
502 fn test_same_seed_same_sequence() {
503 let dist = NoiseDistribution::Gaussian {
504 mean: 0.0f64,
505 std_dev: 0.01,
506 };
507 let mut a =
508 NoiseInjectionScheduler::new_seeded(ConstantScheduler::new(0.1), dist, 0.001, 12345);
509 let mut b =
510 NoiseInjectionScheduler::new(ConstantScheduler::new(0.1), dist, 0.001).with_seed(12345);
511
512 assert_eq!(a.get_learning_rate(), b.get_learning_rate());
513 let seq_a: Vec<f64> = (0..64).map(|_| a.step()).collect();
514 let seq_b: Vec<f64> = (0..64).map(|_| b.step()).collect();
515 assert_eq!(seq_a, seq_b);
516
517 let mut c =
519 NoiseInjectionScheduler::new_seeded(ConstantScheduler::new(0.1), dist, 0.001, 999);
520 let seq_c: Vec<f64> = (0..64).map(|_| c.step()).collect();
521 assert_ne!(seq_a, seq_c);
522 }
523
524 #[test]
525 fn test_reset_restores_deterministic_stream() {
526 let dist = NoiseDistribution::Uniform {
527 min: -0.02f64,
528 max: 0.02,
529 };
530 let mut scheduler =
531 NoiseInjectionScheduler::new_seeded(ConstantScheduler::new(0.1), dist, 0.001, 77);
532
533 let first: Vec<f64> = (0..16).map(|_| scheduler.step()).collect();
534 scheduler.reset();
535 let second: Vec<f64> = (0..16).map(|_| scheduler.step()).collect();
536 assert_eq!(first, second);
537 }
538
539 #[test]
540 fn test_degenerate_distributions_are_finite() {
541 let mut uniform = NoiseInjectionScheduler::new_seeded(
543 ConstantScheduler::new(0.1f64),
544 NoiseDistribution::Uniform { min: 0.0, max: 0.0 },
545 0.001,
546 1,
547 );
548 for _ in 0..10 {
549 assert!(uniform.step().is_finite());
550 }
551
552 let mut cyclical = NoiseInjectionScheduler::new_seeded(
554 ConstantScheduler::new(0.1f64),
555 NoiseDistribution::Cyclical {
556 amplitude: 0.01,
557 period: 0,
558 },
559 0.001,
560 2,
561 );
562 for _ in 0..10 {
563 assert!(cyclical.step().is_finite());
564 }
565
566 let mut decaying = NoiseInjectionScheduler::new_seeded(
567 ConstantScheduler::new(0.1f64),
568 NoiseDistribution::Decaying {
569 initial_scale: 0.05,
570 final_scale: 0.001,
571 decay_steps: 0,
572 },
573 0.001,
574 3,
575 );
576 for _ in 0..10 {
577 assert!(decaying.step().is_finite());
578 }
579 }
580}