use scirs2_core::ndarray::ScalarOperand;
use scirs2_core::numeric::Float;
use std::fmt::Debug;
use crate::schedulers::LearningRateScheduler;
#[derive(Debug, Clone)]
pub struct ReduceOnPlateau<A: Float + Debug> {
current_lr: A,
factor: A,
patience: usize,
min_lr: A,
stagnation_count: usize,
best_metric: Option<A>,
threshold: A,
mode_is_min: bool,
cooldown: usize,
cooldown_counter: usize,
}
impl<A: Float + Debug + Send + Sync> ReduceOnPlateau<A> {
pub fn new(initial_lr: A, factor: A, patience: usize, min_lr: A) -> Self {
Self {
current_lr: initial_lr,
factor,
patience,
min_lr,
stagnation_count: 0,
best_metric: None,
threshold: A::from(1e-4).unwrap_or_else(A::zero),
mode_is_min: true,
cooldown: 0,
cooldown_counter: 0,
}
}
pub fn mode_min(&mut self) -> &mut Self {
self.mode_is_min = true;
self
}
pub fn mode_max(&mut self) -> &mut Self {
self.mode_is_min = false;
self
}
pub fn set_threshold(&mut self, threshold: A) -> &mut Self {
self.threshold = threshold;
self
}
pub fn with_cooldown(mut self, cooldown: usize) -> Self {
self.cooldown = cooldown;
self
}
pub fn set_cooldown(&mut self, cooldown: usize) -> &mut Self {
self.cooldown = cooldown;
self
}
pub fn cooldown(&self) -> usize {
self.cooldown
}
pub fn cooldown_counter(&self) -> usize {
self.cooldown_counter
}
pub fn stagnation_count(&self) -> usize {
self.stagnation_count
}
pub fn best_metric(&self) -> Option<A> {
self.best_metric
}
pub fn step_with_metric(&mut self, metric: A) -> A {
self.update_with_metric(metric)
}
fn update_with_metric(&mut self, metric: A) -> A {
let is_improvement = match self.best_metric {
None => true, Some(best) => {
if self.mode_is_min {
metric < best * (A::one() - self.threshold)
} else {
metric > best * (A::one() + self.threshold)
}
}
};
if is_improvement {
self.best_metric = Some(metric);
self.stagnation_count = 0;
}
if self.cooldown_counter > 0 {
self.cooldown_counter -= 1;
self.stagnation_count = 0;
return self.current_lr;
}
if !is_improvement {
self.stagnation_count += 1;
if self.stagnation_count >= self.patience {
self.current_lr = (self.current_lr * self.factor).max(self.min_lr);
self.stagnation_count = 0;
self.cooldown_counter = self.cooldown;
}
}
self.current_lr
}
}
impl<A: Float + Debug + ScalarOperand + Send + Sync> LearningRateScheduler<A>
for ReduceOnPlateau<A>
{
fn get_learning_rate(&self) -> A {
self.current_lr
}
fn step(&mut self) -> A {
self.current_lr
}
fn step_with_metric(&mut self, metric: A) -> A {
self.update_with_metric(metric)
}
fn reset(&mut self) {
self.stagnation_count = 0;
self.best_metric = None;
self.cooldown_counter = 0;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_reduces_after_patience() {
let mut scheduler = ReduceOnPlateau::new(0.1f64, 0.5, 2, 1e-6);
assert!((scheduler.step_with_metric(1.0) - 0.1).abs() < 1e-12);
assert!((scheduler.step_with_metric(0.9) - 0.1).abs() < 1e-12);
assert!((scheduler.step_with_metric(0.85) - 0.1).abs() < 1e-12);
assert!((scheduler.step_with_metric(0.84) - 0.1).abs() < 1e-12);
assert!((scheduler.step_with_metric(0.84) - 0.1).abs() < 1e-12);
assert!((scheduler.step_with_metric(0.84) - 0.05).abs() < 1e-12);
}
#[test]
fn test_plain_step_is_inert() {
let mut scheduler = ReduceOnPlateau::new(0.1f64, 0.1, 2, 1e-6);
for _ in 0..50 {
assert!((scheduler.step() - 0.1).abs() < 1e-12);
}
assert_eq!(scheduler.stagnation_count(), 0);
assert!(scheduler.best_metric().is_none());
}
#[test]
fn test_trait_object_drives_plateau() {
let mut scheduler: Box<dyn LearningRateScheduler<f64>> =
Box::new(ReduceOnPlateau::new(0.1f64, 0.1, 2, 1e-6));
scheduler.step_with_metric(1.0);
for _ in 0..4 {
scheduler.step_with_metric(1.0);
}
assert!(scheduler.get_learning_rate() < 0.1);
}
#[test]
fn test_cooldown_suppresses_consecutive_reductions() {
let mut without = ReduceOnPlateau::new(1.0f64, 0.5, 1, 1e-9);
without.step_with_metric(1.0);
let mut lrs_without = Vec::new();
for _ in 0..4 {
lrs_without.push(without.step_with_metric(1.0));
}
assert!((lrs_without[0] - 0.5).abs() < 1e-12);
assert!((lrs_without[1] - 0.25).abs() < 1e-12);
assert!((lrs_without[2] - 0.125).abs() < 1e-12);
assert!((lrs_without[3] - 0.0625).abs() < 1e-12);
let mut with = ReduceOnPlateau::new(1.0f64, 0.5, 1, 1e-9).with_cooldown(2);
with.step_with_metric(1.0);
let mut lrs_with = Vec::new();
for _ in 0..4 {
lrs_with.push(with.step_with_metric(1.0));
}
assert!((lrs_with[0] - 0.5).abs() < 1e-12);
assert!((lrs_with[1] - 0.5).abs() < 1e-12);
assert!((lrs_with[2] - 0.5).abs() < 1e-12);
assert!((lrs_with[3] - 0.25).abs() < 1e-12);
assert_eq!(with.cooldown(), 2);
}
#[test]
fn test_min_lr_is_respected() {
let mut scheduler = ReduceOnPlateau::new(1.0f64, 0.1, 1, 0.5);
scheduler.step_with_metric(1.0);
for _ in 0..10 {
scheduler.step_with_metric(1.0);
}
assert!((scheduler.get_learning_rate() - 0.5).abs() < 1e-12);
}
#[test]
fn test_max_mode() {
let mut scheduler = ReduceOnPlateau::new(0.1f64, 0.5, 2, 1e-6);
scheduler.mode_max();
scheduler.step_with_metric(0.5);
scheduler.step_with_metric(0.9);
assert!((scheduler.get_learning_rate() - 0.1).abs() < 1e-12);
scheduler.step_with_metric(0.9);
scheduler.step_with_metric(0.9);
assert!((scheduler.get_learning_rate() - 0.05).abs() < 1e-12);
}
#[test]
fn test_reset_clears_cooldown_and_history() {
let mut scheduler = ReduceOnPlateau::new(1.0f64, 0.5, 1, 1e-9).with_cooldown(3);
scheduler.step_with_metric(1.0);
scheduler.step_with_metric(1.0);
assert!(scheduler.cooldown_counter() > 0);
scheduler.reset();
assert_eq!(scheduler.cooldown_counter(), 0);
assert_eq!(scheduler.stagnation_count(), 0);
assert!(scheduler.best_metric().is_none());
}
}