//! `OneCycle` learning rate scheduler //! //! Implements the one-cycle learning rate policy with three-phase scheduling //! and optional momentum scheduling support. use crate::schedulers::LearningRateScheduler; use crate::{Result, TransformerError}; use serde::{Deserialize, Serialize}; use tracing::{debug, trace}; /// `OneCycle` learning rate scheduler /// /// Implements the one-cycle learning rate policy that consists of three phases: /// 1. Warm-up phase: Linear increase from `base_lr` to `max_lr` /// 2. Cool-down phase: Linear decrease from `max_lr` to `base_lr` /// 3. Annihilation phase: Further decrease from `base_lr` to `final_lr` /// /// # Mathematical Foundation /// /// The scheduler operates in three phases over `total_steps`: /// - Phase 1 (0 to `pct_start` * `total_steps)`: Linear increase /// - Phase 2 (`pct_start` * `total_steps` to (1-final_div_factor_pct) * `total_steps)`: Linear decrease /// - Phase 3 (remaining steps): Further linear decrease to `final_lr` /// /// # Benefits /// - Fast convergence through super-convergence /// - Enables training with very high learning rates /// - Reduces training time significantly /// - Widely used in modern deep learning training /// - Can include momentum scheduling for better optimization #[derive(Debug, Clone, Serialize, Deserialize)] pub struct OneCycleScheduler { /// Base learning rate (starting point) base_lr: f64, /// Maximum learning rate (peak of the cycle) max_lr: f64, /// Total number of training steps total_steps: usize, /// Current step count current_step: usize, /// Percentage of cycle spent in warm-up phase pct_start: f64, /// Percentage of cycle for final decrease phase final_div_factor_pct: f64, /// Final learning rate divisor (`final_lr` = `max_lr` / `final_div_factor`) final_div_factor: f64, /// Whether to include momentum scheduling use_momentum: bool, /// Base momentum (when `use_momentum` is true) base_momentum: Option, /// Maximum momentum (when `use_momentum` is true) max_momentum: Option, } impl OneCycleScheduler { /// Create a new `OneCycle` scheduler with default parameters /// /// # Arguments /// * `base_lr` - Base learning rate (must be positive) /// * `max_lr` - Maximum learning rate (must be > `base_lr`) /// * `total_steps` - Total number of training steps (must be positive) /// /// # Errors /// Returns error if parameters are invalid pub fn new(base_lr: f64, max_lr: f64, total_steps: usize) -> Result { Self::with_params(base_lr, max_lr, total_steps, 0.3, 0.1, 10000.0) } /// Create a new `OneCycle` scheduler with custom parameters /// /// # Arguments /// * `base_lr` - Base learning rate (must be positive) /// * `max_lr` - Maximum learning rate (must be > `base_lr`) /// * `total_steps` - Total number of training steps (must be positive) /// * `pct_start` - Percentage of cycle for warm-up (0.0 to 1.0) /// * `final_div_factor_pct` - Percentage of cycle for final decrease (0.0 to 1.0) /// * `final_div_factor` - Divisor for final learning rate (must be positive) /// /// # Errors /// Returns error if parameters are invalid pub fn with_params( base_lr: f64, max_lr: f64, total_steps: usize, pct_start: f64, final_div_factor_pct: f64, final_div_factor: f64, ) -> Result { if base_lr <= 0.0 { return Err(TransformerError::scheduling(format!( "base_lr {base_lr} must be positive" ))); } if max_lr <= 0.0 { return Err(TransformerError::scheduling(format!( "max_lr {max_lr} must be positive" ))); } if max_lr <= base_lr { return Err(TransformerError::scheduling(format!( "max_lr {max_lr} must be > base_lr {base_lr}" ))); } if total_steps == 0 { return Err(TransformerError::scheduling(format!( "total_steps {total_steps} must be positive" ))); } if pct_start <= 0.0 || pct_start >= 1.0 { return Err(TransformerError::scheduling(format!( "pct_start {pct_start} must be in (0.0, 1.0)" ))); } if final_div_factor_pct <= 0.0 || final_div_factor_pct >= 1.0 { return Err(TransformerError::scheduling(format!( "final_div_factor_pct {final_div_factor_pct} must be in (0.0, 1.0)" ))); } if final_div_factor <= 0.0 { return Err(TransformerError::scheduling(format!( "final_div_factor {final_div_factor} must be positive" ))); } debug!( "Creating OneCycle scheduler: base_lr={}, max_lr={}, total_steps={}, pct_start={}, final_div_factor_pct={}, final_div_factor={}", base_lr, max_lr, total_steps, pct_start, final_div_factor_pct, final_div_factor ); Ok(Self { base_lr, max_lr, total_steps, current_step: 0, pct_start, final_div_factor_pct, final_div_factor, use_momentum: false, base_momentum: None, max_momentum: None, }) } /// Create a new `OneCycle` scheduler with momentum scheduling /// /// # Arguments /// * `base_lr` - Base learning rate (must be positive) /// * `max_lr` - Maximum learning rate (must be > `base_lr`) /// * `total_steps` - Total number of training steps (must be positive) /// * `base_momentum` - Base momentum value (typically 0.85) /// * `max_momentum` - Maximum momentum value (typically 0.95) /// /// # Errors /// Returns error if parameters are invalid pub fn with_momentum( base_lr: f64, max_lr: f64, total_steps: usize, base_momentum: f64, max_momentum: f64, ) -> Result { let mut scheduler = Self::new(base_lr, max_lr, total_steps)?; scheduler.use_momentum = true; scheduler.base_momentum = Some(base_momentum); scheduler.max_momentum = Some(max_momentum); Ok(scheduler) } /// Get the maximum learning rate #[must_use] pub fn max_lr(&self) -> f64 { self.max_lr } /// Get the total number of steps #[must_use] pub fn total_steps(&self) -> usize { self.total_steps } /// Get the percentage of cycle for warm-up #[must_use] pub fn pct_start(&self) -> f64 { self.pct_start } /// Get the percentage of cycle for final decrease #[must_use] pub fn final_div_factor_pct(&self) -> f64 { self.final_div_factor_pct } /// Get the final learning rate divisor #[must_use] pub fn final_div_factor(&self) -> f64 { self.final_div_factor } /// Check if using momentum scheduling #[must_use] pub fn use_momentum(&self) -> bool { self.use_momentum } /// Get the base momentum (if momentum scheduling is enabled) #[must_use] pub fn base_momentum(&self) -> Option { self.base_momentum } /// Get the maximum momentum (if momentum scheduling is enabled) #[must_use] pub fn max_momentum(&self) -> Option { self.max_momentum } /// Get the current phase (0: warm-up, 1: cool-down, 2: annihilation) #[must_use] pub fn get_phase(&self, step: usize) -> usize { let warmup_steps = (self.pct_start * self.total_steps as f64) as usize; let cooldown_end_steps = ((1.0 - self.final_div_factor_pct) * self.total_steps as f64) as usize; if step < warmup_steps { 0 // Warm-up phase } else if step < cooldown_end_steps { 1 // Cool-down phase } else { 2 // Annihilation phase } } /// Get the progress within the current phase (0.0 to 1.0) #[must_use] pub fn phase_progress(&self, step: usize) -> f64 { let warmup_steps = (self.pct_start * self.total_steps as f64) as usize; let cooldown_end_steps = ((1.0 - self.final_div_factor_pct) * self.total_steps as f64) as usize; match self.get_phase(step) { 0 => { // Warm-up phase if warmup_steps == 0 { 1.0 } else { step as f64 / warmup_steps as f64 } } 1 => { // Cool-down phase let cooldown_steps = cooldown_end_steps - warmup_steps; if cooldown_steps == 0 { 1.0 } else { (step - warmup_steps) as f64 / cooldown_steps as f64 } } 2 => { // Annihilation phase let annihilation_steps = self.total_steps - cooldown_end_steps; if annihilation_steps == 0 { 1.0 } else { (step - cooldown_end_steps) as f64 / annihilation_steps as f64 } } _ => unreachable!(), } } /// Get the momentum value for the given step (if momentum scheduling is enabled) #[must_use] pub fn get_momentum(&self, step: usize) -> Option { if !self.use_momentum { return None; } let base_mom = self.base_momentum?; let max_mom = self.max_momentum?; let phase = self.get_phase(step); let progress = self.phase_progress(step); let momentum = match phase { 0 => { // Warm-up: momentum decreases from max to base (inverse of LR) max_mom - (max_mom - base_mom) * progress } 1 => { // Cool-down: momentum increases from base back to max (inverse of LR) base_mom + (max_mom - base_mom) * progress } 2 => { // Annihilation: momentum stays at max max_mom } _ => unreachable!(), }; Some(momentum) } } impl LearningRateScheduler for OneCycleScheduler { fn get_lr(&self, _epoch: usize, step: usize) -> f64 { let step = step.min(self.total_steps); let phase = self.get_phase(step); let progress = self.phase_progress(step); let lr = match phase { 0 => { // Warm-up phase: linear increase from base_lr to max_lr self.base_lr + (self.max_lr - self.base_lr) * progress } 1 => { // Cool-down phase: linear decrease from max_lr to base_lr self.max_lr - (self.max_lr - self.base_lr) * progress } 2 => { // Annihilation phase: linear decrease from base_lr to final_lr let final_lr = self.max_lr / self.final_div_factor; self.base_lr - (self.base_lr - final_lr) * progress } _ => unreachable!(), }; trace!( "OneCycle step {}: phase={}, progress={:.4}, lr={:.6}", step, phase, progress, lr ); lr } fn step(&mut self) { self.current_step += 1; trace!("OneCycle scheduler stepped to: {}", self.current_step); } fn current_step(&self) -> usize { self.current_step } fn reset(&mut self) { self.current_step = 0; debug!("Reset OneCycle scheduler"); } fn scheduler_type(&self) -> &'static str { "OneCycle" } fn base_lr(&self) -> f64 { self.base_lr } } #[cfg(all(test, feature = "disabled_tests"))] mod tests { use super::*; #[test] fn test_one_cycle_creation() { let scheduler = OneCycleScheduler::new(0.001, 0.01, 1000).unwrap(); assert_eq!(scheduler.base_lr(), 0.001); assert_eq!(scheduler.max_lr(), 0.01); assert_eq!(scheduler.total_steps(), 1000); assert_eq!(scheduler.current_step(), 0); assert_eq!(scheduler.pct_start(), 0.3); // Default assert_eq!(scheduler.final_div_factor(), 10000.0); // Default assert!(!scheduler.use_momentum()); } #[test] fn test_one_cycle_with_custom_params() { let scheduler = OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.25, 0.1, 1000.0).unwrap(); assert_eq!(scheduler.base_lr(), 0.001); assert_eq!(scheduler.max_lr(), 0.01); assert_eq!(scheduler.total_steps(), 1000); assert_eq!(scheduler.pct_start(), 0.25); assert_eq!(scheduler.final_div_factor(), 1000.0); assert_eq!(scheduler.final_div_factor_pct(), 0.1); } #[test] fn test_one_cycle_with_momentum() { let scheduler = OneCycleScheduler::with_momentum(0.001, 0.01, 1000, 0.85, 0.95).unwrap(); assert!(scheduler.use_momentum()); assert_eq!(scheduler.base_momentum(), Some(0.85)); assert_eq!(scheduler.max_momentum(), Some(0.95)); } #[test] fn test_one_cycle_invalid_params() { // Negative base_lr assert!(OneCycleScheduler::new(-0.001, 0.01, 1000).is_err()); // Negative max_lr assert!(OneCycleScheduler::new(0.001, -0.01, 1000).is_err()); // max_lr <= base_lr assert!(OneCycleScheduler::new(0.01, 0.001, 1000).is_err()); assert!(OneCycleScheduler::new(0.01, 0.01, 1000).is_err()); // Zero total_steps assert!(OneCycleScheduler::new(0.001, 0.01, 0).is_err()); // Invalid pct_start assert!(OneCycleScheduler::with_params(0.001, 0.01, 1000, -0.1, 0.1, 1000.0).is_err()); assert!(OneCycleScheduler::with_params(0.001, 0.01, 1000, 1.0, 0.1, 1000.0).is_err()); // Invalid final_div_factor_pct assert!(OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.3, -0.1, 1000.0).is_err()); assert!(OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.3, 1.0, 1000.0).is_err()); // Invalid final_div_factor assert!(OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.3, 0.1, 0.0).is_err()); } #[test] fn test_one_cycle_three_phases() { let scheduler = OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.3, 0.1, 100.0).unwrap(); // Phase 1: Warm-up (0 to 300 steps) // At step 0, should be at base_lr let lr_0 = scheduler.get_lr(0, 0); assert!((lr_0 - 0.001).abs() < 1e-10); // At step 150 (middle of warm-up), should be halfway to max_lr let lr_150 = scheduler.get_lr(0, 150); let expected_150 = 0.001 + (0.01 - 0.001) * 0.5; assert!((lr_150 - expected_150).abs() < 1e-6); // At step 300 (end of warm-up), should be at max_lr let lr_300 = scheduler.get_lr(0, 300); assert!((lr_300 - 0.01).abs() < 1e-6); // Phase 2: Cool-down (300 to 900 steps) // At step 600 (middle of cool-down), should be halfway back to base_lr let lr_600 = scheduler.get_lr(0, 600); let expected_600 = 0.01 - (0.01 - 0.001) * 0.5; assert!((lr_600 - expected_600).abs() < 1e-6); // At step 900 (end of cool-down), should be at base_lr let lr_900 = scheduler.get_lr(0, 900); assert!((lr_900 - 0.001).abs() < 1e-6); // Phase 3: Annihilation (900 to 1000 steps) // At step 1000, should be at final_lr = max_lr / final_div_factor let final_lr = 0.01 / 100.0; let lr_1000 = scheduler.get_lr(0, 1000); assert!((lr_1000 - final_lr).abs() < 1e-8); } #[test] fn test_phase_detection() { let scheduler = OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.3, 0.1, 100.0).unwrap(); // Warm-up phase assert_eq!(scheduler.get_phase(0), 0); assert_eq!(scheduler.get_phase(150), 0); assert_eq!(scheduler.get_phase(299), 0); // Cool-down phase assert_eq!(scheduler.get_phase(300), 1); assert_eq!(scheduler.get_phase(600), 1); assert_eq!(scheduler.get_phase(899), 1); // Annihilation phase assert_eq!(scheduler.get_phase(900), 2); assert_eq!(scheduler.get_phase(950), 2); assert_eq!(scheduler.get_phase(1000), 2); } #[test] fn test_momentum_scheduling() { let scheduler = OneCycleScheduler::with_momentum(0.001, 0.01, 1000, 0.85, 0.95).unwrap(); // At step 0, momentum should be at max (inverse of LR) let mom_0 = scheduler.get_momentum(0).unwrap(); assert!((mom_0 - 0.95).abs() < 1e-10); // At step 300 (end of warm-up), momentum should be at base let mom_300 = scheduler.get_momentum(300).unwrap(); assert!((mom_300 - 0.85).abs() < 1e-6); // At step 900 (end of cool-down), momentum should be back to max let mom_900 = scheduler.get_momentum(900).unwrap(); assert!((mom_900 - 0.95).abs() < 1e-6); } #[test] fn test_step_and_reset() { let mut scheduler = OneCycleScheduler::new(0.001, 0.01, 1000).unwrap(); assert_eq!(scheduler.current_step(), 0); scheduler.step(); assert_eq!(scheduler.current_step(), 1); scheduler.step(); scheduler.step(); assert_eq!(scheduler.current_step(), 3); scheduler.reset(); assert_eq!(scheduler.current_step(), 0); } #[test] fn test_scheduler_type() { let scheduler = OneCycleScheduler::new(0.001, 0.01, 1000).unwrap(); assert_eq!(scheduler.scheduler_type(), "OneCycle"); } #[test] fn test_phase_progress() { let scheduler = OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.3, 0.1, 100.0).unwrap(); // Warm-up phase progress assert_eq!(scheduler.phase_progress(0), 0.0); assert_eq!(scheduler.phase_progress(150), 0.5); assert!((scheduler.phase_progress(300) - 1.0).abs() < 1e-10); // Cool-down phase progress assert_eq!(scheduler.phase_progress(300), 0.0); assert_eq!(scheduler.phase_progress(600), 0.5); assert!((scheduler.phase_progress(900) - 1.0).abs() < 1e-10); // Annihilation phase progress assert_eq!(scheduler.phase_progress(900), 0.0); assert_eq!(scheduler.phase_progress(950), 0.5); assert!((scheduler.phase_progress(1000) - 1.0).abs() < 1e-10); } #[test] fn test_edge_cases() { // Very small pct_start let scheduler = OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.01, 0.01, 100.0).unwrap(); // Should still work correctly let lr_5 = scheduler.get_lr(0, 5); // Middle of tiny warm-up phase assert!(lr_5 > 0.001 && lr_5 < 0.01); // Very large final_div_factor let large_div_scheduler = OneCycleScheduler::with_params(0.001, 0.01, 1000, 0.3, 0.1, 1000000.0).unwrap(); let final_lr = large_div_scheduler.get_lr(0, 1000); assert!(final_lr < 1e-8); // Should be very small } }