Consistent formatting pass: line wrapping, import sorting, trailing whitespace removal, let-chain indentation, merged derive attributes, and unsafe block reformatting. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
554 lines
19 KiB
Rust
554 lines
19 KiB
Rust
//! `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<f64>,
|
|
/// Maximum momentum (when `use_momentum` is true)
|
|
max_momentum: Option<f64>,
|
|
}
|
|
|
|
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> {
|
|
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<Self> {
|
|
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<Self> {
|
|
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<f64> {
|
|
self.base_momentum
|
|
}
|
|
|
|
/// Get the maximum momentum (if momentum scheduling is enabled)
|
|
#[must_use]
|
|
pub fn max_momentum(&self) -> Option<f64> {
|
|
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<f64> {
|
|
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
|
|
}
|
|
}
|