1use std::f64::consts::PI;
15
16pub trait ParameterSchedule: Send + Sync {
20 fn value_at(&self, generation: usize, max_generations: usize) -> f64;
22}
23
24#[derive(Clone, Debug)]
26pub struct ConstantSchedule {
27 pub value: f64,
29}
30
31impl ConstantSchedule {
32 pub fn new(value: f64) -> Self {
34 Self { value }
35 }
36}
37
38impl ParameterSchedule for ConstantSchedule {
39 fn value_at(&self, _generation: usize, _max_generations: usize) -> f64 {
40 self.value
41 }
42}
43
44#[derive(Clone, Debug)]
46pub struct LinearAnnealing {
47 pub start: f64,
49 pub end: f64,
51}
52
53impl LinearAnnealing {
54 pub fn new(start: f64, end: f64) -> Self {
56 Self { start, end }
57 }
58
59 pub fn decreasing(start: f64, end: f64) -> Self {
61 Self::new(start, end)
62 }
63
64 pub fn increasing(start: f64, end: f64) -> Self {
66 Self::new(start, end)
67 }
68}
69
70impl ParameterSchedule for LinearAnnealing {
71 fn value_at(&self, generation: usize, max_generations: usize) -> f64 {
72 if max_generations == 0 {
73 return self.start;
74 }
75 let t = generation as f64 / max_generations as f64;
76 self.start + (self.end - self.start) * t
77 }
78}
79
80#[derive(Clone, Debug)]
82pub struct ExponentialDecay {
83 pub initial: f64,
85 pub decay_rate: f64,
87 pub minimum: f64,
89}
90
91impl ExponentialDecay {
92 pub fn new(initial: f64, decay_rate: f64) -> Self {
94 Self {
95 initial,
96 decay_rate,
97 minimum: 0.0,
98 }
99 }
100
101 pub fn with_minimum(mut self, minimum: f64) -> Self {
103 self.minimum = minimum;
104 self
105 }
106}
107
108impl ParameterSchedule for ExponentialDecay {
109 fn value_at(&self, generation: usize, _max_generations: usize) -> f64 {
110 (self.initial * (-self.decay_rate * generation as f64).exp()).max(self.minimum)
111 }
112}
113
114#[derive(Clone, Debug)]
118pub struct CosineAnnealing {
119 pub max_value: f64,
121 pub min_value: f64,
123 pub period: Option<usize>,
125}
126
127impl CosineAnnealing {
128 pub fn new(max_value: f64, min_value: f64) -> Self {
130 Self {
131 max_value,
132 min_value,
133 period: None,
134 }
135 }
136
137 pub fn with_warm_restarts(mut self, period: usize) -> Self {
139 self.period = Some(period);
140 self
141 }
142}
143
144impl ParameterSchedule for CosineAnnealing {
145 fn value_at(&self, generation: usize, max_generations: usize) -> f64 {
146 let effective_gen = match self.period {
147 Some(period) if period > 0 => generation % period,
148 _ => generation,
149 };
150 let effective_max = match self.period {
151 Some(period) if period > 0 => period,
152 _ => max_generations,
153 };
154
155 if effective_max == 0 {
156 return self.max_value;
157 }
158
159 let t = effective_gen as f64 / effective_max as f64;
160 self.min_value + 0.5 * (self.max_value - self.min_value) * (1.0 + (PI * t).cos())
161 }
162}
163
164#[derive(Clone, Debug)]
166pub struct StepSchedule {
167 pub steps: Vec<(usize, f64)>,
169 pub initial: f64,
171}
172
173impl StepSchedule {
174 pub fn new(initial: f64, steps: Vec<(usize, f64)>) -> Self {
176 let mut steps = steps;
177 steps.sort_by_key(|(gen, _)| *gen);
178 Self { steps, initial }
179 }
180
181 pub fn single_step(initial: f64, step_gen: usize, step_value: f64) -> Self {
183 Self::new(initial, vec![(step_gen, step_value)])
184 }
185}
186
187impl ParameterSchedule for StepSchedule {
188 fn value_at(&self, generation: usize, _max_generations: usize) -> f64 {
189 let mut value = self.initial;
190 for &(step_gen, step_value) in &self.steps {
191 if generation >= step_gen {
192 value = step_value;
193 } else {
194 break;
195 }
196 }
197 value
198 }
199}
200
201#[derive(Clone, Debug)]
205pub struct PolynomialDecay {
206 pub initial: f64,
208 pub power: f64,
210 pub minimum: f64,
212}
213
214impl PolynomialDecay {
215 pub fn new(initial: f64, power: f64) -> Self {
217 Self {
218 initial,
219 power,
220 minimum: 0.0,
221 }
222 }
223
224 pub fn with_minimum(mut self, minimum: f64) -> Self {
226 self.minimum = minimum;
227 self
228 }
229}
230
231impl ParameterSchedule for PolynomialDecay {
232 fn value_at(&self, generation: usize, max_generations: usize) -> f64 {
233 if max_generations == 0 {
234 return self.initial;
235 }
236 let t = generation as f64 / max_generations as f64;
237 let decay = (1.0 - t).max(0.0).powf(self.power);
238 self.minimum + (self.initial - self.minimum) * decay
239 }
240}
241
242#[derive(Clone, Debug)]
244pub struct CyclicalSchedule {
245 pub base: f64,
247 pub max_value: f64,
249 pub step_size: usize,
251}
252
253impl CyclicalSchedule {
254 pub fn new(base: f64, max_value: f64, step_size: usize) -> Self {
256 Self {
257 base,
258 max_value,
259 step_size,
260 }
261 }
262}
263
264impl ParameterSchedule for CyclicalSchedule {
265 fn value_at(&self, generation: usize, _max_generations: usize) -> f64 {
266 if self.step_size == 0 {
267 return self.base;
268 }
269
270 let cycle = generation / (2 * self.step_size);
271 let x = (generation as f64 / self.step_size as f64) - 2.0 * cycle as f64;
272 let scale = (1.0 - (x - 1.0).abs()).max(0.0);
273 self.base + (self.max_value - self.base) * scale
274 }
275}
276
277#[derive(Clone, Debug)]
279pub enum DynamicSchedule {
280 Constant(ConstantSchedule),
281 Linear(LinearAnnealing),
282 Exponential(ExponentialDecay),
283 Cosine(CosineAnnealing),
284 Step(StepSchedule),
285 Polynomial(PolynomialDecay),
286 Cyclical(CyclicalSchedule),
287}
288
289impl ParameterSchedule for DynamicSchedule {
290 fn value_at(&self, generation: usize, max_generations: usize) -> f64 {
291 match self {
292 Self::Constant(s) => s.value_at(generation, max_generations),
293 Self::Linear(s) => s.value_at(generation, max_generations),
294 Self::Exponential(s) => s.value_at(generation, max_generations),
295 Self::Cosine(s) => s.value_at(generation, max_generations),
296 Self::Step(s) => s.value_at(generation, max_generations),
297 Self::Polynomial(s) => s.value_at(generation, max_generations),
298 Self::Cyclical(s) => s.value_at(generation, max_generations),
299 }
300 }
301}
302
303impl From<ConstantSchedule> for DynamicSchedule {
304 fn from(s: ConstantSchedule) -> Self {
305 Self::Constant(s)
306 }
307}
308
309impl From<LinearAnnealing> for DynamicSchedule {
310 fn from(s: LinearAnnealing) -> Self {
311 Self::Linear(s)
312 }
313}
314
315impl From<ExponentialDecay> for DynamicSchedule {
316 fn from(s: ExponentialDecay) -> Self {
317 Self::Exponential(s)
318 }
319}
320
321impl From<CosineAnnealing> for DynamicSchedule {
322 fn from(s: CosineAnnealing) -> Self {
323 Self::Cosine(s)
324 }
325}
326
327impl From<StepSchedule> for DynamicSchedule {
328 fn from(s: StepSchedule) -> Self {
329 Self::Step(s)
330 }
331}
332
333impl From<PolynomialDecay> for DynamicSchedule {
334 fn from(s: PolynomialDecay) -> Self {
335 Self::Polynomial(s)
336 }
337}
338
339impl From<CyclicalSchedule> for DynamicSchedule {
340 fn from(s: CyclicalSchedule) -> Self {
341 Self::Cyclical(s)
342 }
343}
344
345#[derive(Clone, Debug)]
347pub struct CompositeSchedule {
348 pub phases: Vec<(usize, DynamicSchedule)>,
350}
351
352impl CompositeSchedule {
353 pub fn new() -> Self {
355 Self { phases: Vec::new() }
356 }
357
358 pub fn add_phase<S: Into<DynamicSchedule>>(mut self, end_gen: usize, schedule: S) -> Self {
360 self.phases.push((end_gen, schedule.into()));
361 self.phases.sort_by_key(|(gen, _)| *gen);
362 self
363 }
364}
365
366impl Default for CompositeSchedule {
367 fn default() -> Self {
368 Self::new()
369 }
370}
371
372impl ParameterSchedule for CompositeSchedule {
373 fn value_at(&self, generation: usize, _max_generations: usize) -> f64 {
374 let mut prev_end = 0;
375 for (end_gen, schedule) in &self.phases {
376 if generation < *end_gen {
377 let phase_duration = end_gen - prev_end;
378 let phase_gen = generation - prev_end;
379 return schedule.value_at(phase_gen, phase_duration);
380 }
381 prev_end = *end_gen;
382 }
383 if let Some((end_gen, schedule)) = self.phases.last() {
385 let phase_duration = end_gen
386 - self
387 .phases
388 .get(self.phases.len().saturating_sub(2))
389 .map(|(e, _)| *e)
390 .unwrap_or(0);
391 schedule.value_at(phase_duration, phase_duration)
392 } else {
393 0.0
394 }
395 }
396}
397
398#[cfg(test)]
399mod tests {
400 use super::*;
401 use approx::assert_relative_eq;
402
403 #[test]
404 fn test_constant_schedule() {
405 let schedule = ConstantSchedule::new(0.5);
406 assert_relative_eq!(schedule.value_at(0, 100), 0.5);
407 assert_relative_eq!(schedule.value_at(50, 100), 0.5);
408 assert_relative_eq!(schedule.value_at(100, 100), 0.5);
409 }
410
411 #[test]
412 fn test_linear_annealing() {
413 let schedule = LinearAnnealing::new(1.0, 0.0);
414 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
415 assert_relative_eq!(schedule.value_at(50, 100), 0.5);
416 assert_relative_eq!(schedule.value_at(100, 100), 0.0);
417 }
418
419 #[test]
420 fn test_linear_annealing_increasing() {
421 let schedule = LinearAnnealing::increasing(0.1, 0.9);
422 assert_relative_eq!(schedule.value_at(0, 100), 0.1);
423 assert_relative_eq!(schedule.value_at(100, 100), 0.9);
424 }
425
426 #[test]
427 fn test_exponential_decay() {
428 let schedule = ExponentialDecay::new(1.0, 0.1);
429 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
430 assert!(schedule.value_at(10, 100) < 1.0);
431 assert!(schedule.value_at(50, 100) < schedule.value_at(10, 100));
432 }
433
434 #[test]
435 fn test_exponential_decay_with_minimum() {
436 let schedule = ExponentialDecay::new(1.0, 0.1).with_minimum(0.1);
437 assert!(schedule.value_at(1000, 100) >= 0.1);
438 }
439
440 #[test]
441 fn test_cosine_annealing() {
442 let schedule = CosineAnnealing::new(1.0, 0.0);
443 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
444 assert_relative_eq!(schedule.value_at(100, 100), 0.0, epsilon = 1e-10);
445 assert_relative_eq!(schedule.value_at(50, 100), 0.5, epsilon = 1e-10);
447 }
448
449 #[test]
450 fn test_cosine_annealing_warm_restarts() {
451 let schedule = CosineAnnealing::new(1.0, 0.0).with_warm_restarts(50);
452 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
453 assert_relative_eq!(schedule.value_at(50, 100), 1.0); assert_relative_eq!(schedule.value_at(25, 100), 0.5, epsilon = 1e-10);
455 }
456
457 #[test]
458 fn test_step_schedule() {
459 let schedule = StepSchedule::new(1.0, vec![(25, 0.5), (75, 0.1)]);
460 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
461 assert_relative_eq!(schedule.value_at(24, 100), 1.0);
462 assert_relative_eq!(schedule.value_at(25, 100), 0.5);
463 assert_relative_eq!(schedule.value_at(74, 100), 0.5);
464 assert_relative_eq!(schedule.value_at(75, 100), 0.1);
465 }
466
467 #[test]
468 fn test_polynomial_decay() {
469 let schedule = PolynomialDecay::new(1.0, 2.0).with_minimum(0.0);
470 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
471 assert_relative_eq!(schedule.value_at(100, 100), 0.0);
472 assert_relative_eq!(schedule.value_at(50, 100), 0.25);
474 }
475
476 #[test]
481 fn test_polynomial_decay_no_offset_at_start() {
482 let schedule = PolynomialDecay::new(1.0, 2.0).with_minimum(0.2);
483 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
485 assert_relative_eq!(schedule.value_at(100, 100), 0.2);
487 assert_relative_eq!(schedule.value_at(50, 100), 0.4);
489 }
490
491 #[test]
492 fn test_cyclical_schedule() {
493 let schedule = CyclicalSchedule::new(0.0, 1.0, 10);
494 assert_relative_eq!(schedule.value_at(0, 100), 0.0);
495 assert_relative_eq!(schedule.value_at(10, 100), 1.0);
496 assert_relative_eq!(schedule.value_at(20, 100), 0.0);
497 assert_relative_eq!(schedule.value_at(30, 100), 1.0);
498 }
499
500 #[test]
501 fn test_linear_annealing_decreasing() {
502 let schedule = LinearAnnealing::decreasing(0.9, 0.1);
503 assert_relative_eq!(schedule.value_at(0, 100), 0.9);
504 assert_relative_eq!(schedule.value_at(100, 100), 0.1);
505 }
506
507 #[test]
508 fn test_linear_annealing_zero_max_generations() {
509 let schedule = LinearAnnealing::new(1.0, 0.0);
510 assert_relative_eq!(schedule.value_at(0, 0), 1.0);
511 }
512
513 #[test]
514 fn test_step_schedule_single_step() {
515 let schedule = StepSchedule::single_step(1.0, 50, 0.5);
516 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
517 assert_relative_eq!(schedule.value_at(49, 100), 1.0);
518 assert_relative_eq!(schedule.value_at(50, 100), 0.5);
519 assert_relative_eq!(schedule.value_at(100, 100), 0.5);
520 }
521
522 #[test]
523 fn test_polynomial_decay_zero_max_generations() {
524 let schedule = PolynomialDecay::new(1.0, 2.0);
525 assert_relative_eq!(schedule.value_at(0, 0), 1.0);
526 }
527
528 #[test]
529 fn test_cyclical_schedule_zero_step_size() {
530 let schedule = CyclicalSchedule::new(0.5, 1.0, 0);
531 assert_relative_eq!(schedule.value_at(0, 100), 0.5);
532 assert_relative_eq!(schedule.value_at(50, 100), 0.5);
533 }
534
535 #[test]
536 fn test_cosine_annealing_zero_max_generations() {
537 let schedule = CosineAnnealing::new(1.0, 0.0);
538 assert_relative_eq!(schedule.value_at(0, 0), 1.0);
539 }
540
541 #[test]
542 fn test_cosine_annealing_warm_restarts_zero_period() {
543 let schedule = CosineAnnealing::new(1.0, 0.0).with_warm_restarts(0);
544 assert_relative_eq!(schedule.value_at(50, 100), 0.5, epsilon = 1e-10);
546 }
547
548 #[test]
549 fn test_dynamic_schedule_from_conversions() {
550 let constant: DynamicSchedule = ConstantSchedule::new(0.5).into();
551 assert_relative_eq!(constant.value_at(50, 100), 0.5);
552
553 let linear: DynamicSchedule = LinearAnnealing::new(1.0, 0.0).into();
554 assert_relative_eq!(linear.value_at(50, 100), 0.5);
555
556 let exponential: DynamicSchedule = ExponentialDecay::new(1.0, 0.1).into();
557 assert!(exponential.value_at(10, 100) < 1.0);
558
559 let cosine: DynamicSchedule = CosineAnnealing::new(1.0, 0.0).into();
560 assert_relative_eq!(cosine.value_at(50, 100), 0.5, epsilon = 1e-10);
561
562 let step: DynamicSchedule = StepSchedule::new(1.0, vec![(50, 0.5)]).into();
563 assert_relative_eq!(step.value_at(50, 100), 0.5);
564
565 let polynomial: DynamicSchedule = PolynomialDecay::new(1.0, 2.0).into();
566 assert_relative_eq!(polynomial.value_at(50, 100), 0.25);
567
568 let cyclical: DynamicSchedule = CyclicalSchedule::new(0.0, 1.0, 10).into();
569 assert_relative_eq!(cyclical.value_at(10, 100), 1.0);
570 }
571
572 #[test]
573 fn test_composite_schedule() {
574 let schedule = CompositeSchedule::new()
575 .add_phase(50, ConstantSchedule::new(1.0))
576 .add_phase(100, LinearAnnealing::new(1.0, 0.0));
577
578 assert_relative_eq!(schedule.value_at(0, 100), 1.0);
580 assert_relative_eq!(schedule.value_at(25, 100), 1.0);
581
582 assert_relative_eq!(schedule.value_at(50, 100), 1.0);
584 assert_relative_eq!(schedule.value_at(75, 100), 0.5);
585 }
586
587 #[test]
588 fn test_composite_schedule_empty() {
589 let schedule = CompositeSchedule::default();
590 assert_relative_eq!(schedule.value_at(50, 100), 0.0);
591 }
592
593 #[test]
594 fn test_composite_schedule_past_all_phases() {
595 let schedule = CompositeSchedule::new().add_phase(50, ConstantSchedule::new(0.5));
596
597 assert_relative_eq!(schedule.value_at(100, 100), 0.5);
599 }
600}