use scirs2_core::ndarray::ScalarOperand;
use scirs2_core::numeric::Float;
use std::fmt::Debug;
use crate::error::{OptimError, Result};
use crate::schedulers::LearningRateScheduler;
fn from_f64<A: Float>(v: f64) -> A {
A::from(v).unwrap_or_else(A::zero)
}
fn from_usize<A: Float>(v: usize) -> A {
A::from(v).unwrap_or_else(A::zero)
}
fn denom_from_usize<A: Float>(v: usize) -> A {
match A::from(v) {
Some(x) if x != A::zero() => x,
_ => A::one(),
}
}
#[derive(Debug, Clone)]
pub struct CosineAnnealing<A: Float + Debug> {
initial_lr: A,
min_lr: A,
t_max: usize,
warm_restart: bool,
step: usize,
current_lr: A,
}
impl<A: Float + Debug + Send + Sync> CosineAnnealing<A> {
pub fn new(initial_lr: A, min_lr: A, t_max: usize, warm_restart: bool) -> Self {
Self {
initial_lr,
min_lr,
t_max: t_max.max(1),
warm_restart,
step: 0,
current_lr: initial_lr,
}
}
pub fn try_new(initial_lr: A, min_lr: A, t_max: usize, warm_restart: bool) -> Result<Self> {
if t_max == 0 {
return Err(OptimError::InvalidConfig(
"CosineAnnealing requires t_max > 0".to_string(),
));
}
if !initial_lr.is_finite() || !min_lr.is_finite() {
return Err(OptimError::InvalidConfig(
"CosineAnnealing requires finite initial_lr and min_lr".to_string(),
));
}
if min_lr < A::zero() {
return Err(OptimError::InvalidConfig(
"CosineAnnealing requires min_lr >= 0".to_string(),
));
}
if min_lr > initial_lr {
return Err(OptimError::InvalidConfig(
"CosineAnnealing requires min_lr <= initial_lr".to_string(),
));
}
Ok(Self::new(initial_lr, min_lr, t_max, warm_restart))
}
pub fn t_max(&self) -> usize {
self.t_max
}
pub fn warm_restart(&self) -> bool {
self.warm_restart
}
fn cosine_lr(&self, t_cur: usize) -> A {
let pi = from_f64::<A>(std::f64::consts::PI);
let half = from_f64::<A>(0.5);
let progress = from_usize::<A>(t_cur) / denom_from_usize::<A>(self.t_max);
let cos_term = A::one() + (pi * progress).cos();
self.min_lr + half * (self.initial_lr - self.min_lr) * cos_term
}
}
impl<A: Float + Debug + ScalarOperand + Send + Sync> LearningRateScheduler<A>
for CosineAnnealing<A>
{
fn get_learning_rate(&self) -> A {
self.current_lr
}
fn step(&mut self) -> A {
self.step += 1;
let t_cur = if self.warm_restart {
if self.step >= self.t_max {
self.step = 0;
}
self.step
} else {
if self.step > self.t_max {
self.step = self.t_max;
}
self.step
};
self.current_lr = self.cosine_lr(t_cur);
self.current_lr
}
fn reset(&mut self) {
self.step = 0;
self.current_lr = self.initial_lr;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_no_restart_anneals_once_and_pins_at_min() {
let mut scheduler = CosineAnnealing::new(0.1f64, 0.001, 10, false);
let mut previous = scheduler.get_learning_rate();
assert!((previous - 0.1).abs() < 1e-12);
for _ in 0..10 {
let lr = scheduler.step();
assert!(lr <= previous + 1e-12);
previous = lr;
}
assert!((scheduler.get_learning_rate() - 0.001).abs() < 1e-12);
for _ in 0..25 {
assert!((scheduler.step() - 0.001).abs() < 1e-12);
}
}
#[test]
fn test_warm_restart_restarts() {
let mut scheduler = CosineAnnealing::new(0.1f64, 0.001, 10, true);
let mut lrs = Vec::new();
for _ in 0..20 {
lrs.push(scheduler.step());
}
assert!((lrs[9] - 0.1).abs() < 1e-12);
assert!(lrs[9] > lrs[8]);
assert!((lrs[19] - 0.1).abs() < 1e-12);
}
#[test]
fn test_zero_t_max_is_clamped_not_nan() {
let mut scheduler = CosineAnnealing::new(0.1f64, 0.001, 0, false);
assert_eq!(scheduler.t_max(), 1);
for _ in 0..5 {
assert!(scheduler.step().is_finite());
}
}
#[test]
fn test_try_new_rejects_zero_t_max() {
assert!(CosineAnnealing::try_new(0.1f64, 0.001, 0, false).is_err());
assert!(CosineAnnealing::try_new(0.1f64, 0.2, 10, false).is_err());
assert!(CosineAnnealing::try_new(0.1f64, -1.0, 10, false).is_err());
assert!(CosineAnnealing::try_new(f64::NAN, 0.0, 10, false).is_err());
assert!(CosineAnnealing::try_new(0.1f64, 0.001, 10, false).is_ok());
}
#[test]
fn test_reset() {
let mut scheduler = CosineAnnealing::new(0.1f64, 0.001, 10, false);
for _ in 0..20 {
scheduler.step();
}
scheduler.reset();
assert!((scheduler.get_learning_rate() - 0.1).abs() < 1e-12);
for _ in 0..10 {
scheduler.step();
}
assert!((scheduler.get_learning_rate() - 0.001).abs() < 1e-12);
}
}