Initial commit
This commit is contained in:
@@ -0,0 +1,195 @@
|
||||
//! Warmup learning rate scheduler
|
||||
//!
|
||||
//! Implements linear warmup for learning rate, essential for stable transformer training.
|
||||
|
||||
use crate::schedulers::LearningRateScheduler;
|
||||
use crate::{Result, TransformerError};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, trace};
|
||||
|
||||
/// Linear warmup learning rate scheduler
|
||||
///
|
||||
/// During the warmup phase, the learning rate increases linearly from 0 to the base learning rate.
|
||||
/// This helps stabilize training in the early stages, especially for large models.
|
||||
///
|
||||
/// # Mathematical Foundation
|
||||
///
|
||||
/// For step t during warmup (t < `warmup_steps)`:
|
||||
/// - lr(t) = `base_lr` * (t / `warmup_steps`)
|
||||
///
|
||||
/// After warmup (t >= `warmup_steps)`:
|
||||
/// - lr(t) = `base_lr`
|
||||
///
|
||||
/// # Benefits
|
||||
/// - Prevents gradient explosion in early training
|
||||
/// - Enables training with larger learning rates
|
||||
/// - Essential for transformer model stability
|
||||
/// - Commonly used in BERT, GPT, and other transformers
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct WarmupScheduler {
|
||||
/// Target learning rate after warmup
|
||||
base_lr: f64,
|
||||
/// Number of warmup steps
|
||||
warmup_steps: usize,
|
||||
/// Current step count
|
||||
current_step: usize,
|
||||
}
|
||||
|
||||
impl WarmupScheduler {
|
||||
/// Create a new warmup scheduler
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `base_lr` - Target learning rate after warmup (must be positive)
|
||||
/// * `warmup_steps` - Number of steps for warmup (must be positive)
|
||||
///
|
||||
/// # Errors
|
||||
/// Returns error if parameters are invalid
|
||||
pub fn new(base_lr: f64, warmup_steps: usize) -> Result<Self> {
|
||||
if base_lr <= 0.0 {
|
||||
return Err(TransformerError::generic(format!(
|
||||
"base_lr {base_lr} must be positive"
|
||||
)));
|
||||
}
|
||||
|
||||
if warmup_steps == 0 {
|
||||
return Err(TransformerError::generic(format!(
|
||||
"warmup_steps {warmup_steps} must be greater than 0"
|
||||
)));
|
||||
}
|
||||
|
||||
debug!(
|
||||
"Creating warmup scheduler: base_lr={}, warmup_steps={}",
|
||||
base_lr, warmup_steps
|
||||
);
|
||||
|
||||
Ok(Self {
|
||||
base_lr,
|
||||
warmup_steps,
|
||||
current_step: 0,
|
||||
})
|
||||
}
|
||||
|
||||
/// Get the number of warmup steps
|
||||
#[must_use]
|
||||
pub fn warmup_steps(&self) -> usize {
|
||||
self.warmup_steps
|
||||
}
|
||||
|
||||
/// Check if currently in warmup phase
|
||||
#[must_use]
|
||||
pub fn in_warmup(&self, step: usize) -> bool {
|
||||
step < self.warmup_steps
|
||||
}
|
||||
|
||||
/// Get the warmup progress (0.0 to 1.0)
|
||||
#[must_use]
|
||||
pub fn warmup_progress(&self, step: usize) -> f64 {
|
||||
if step >= self.warmup_steps {
|
||||
1.0
|
||||
} else {
|
||||
step as f64 / self.warmup_steps as f64
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LearningRateScheduler for WarmupScheduler {
|
||||
fn get_lr(&self, _epoch: usize, step: usize) -> f64 {
|
||||
if step < self.warmup_steps {
|
||||
// Linear warmup: lr = base_lr * (step / warmup_steps)
|
||||
let progress = step as f64 / self.warmup_steps as f64;
|
||||
let lr = self.base_lr * progress;
|
||||
trace!(
|
||||
"Warmup step {}: progress={:.4}, lr={:.6}",
|
||||
step, progress, lr
|
||||
);
|
||||
lr
|
||||
} else {
|
||||
// After warmup: lr = base_lr
|
||||
trace!("Post-warmup step {}: lr={:.6}", step, self.base_lr);
|
||||
self.base_lr
|
||||
}
|
||||
}
|
||||
|
||||
fn step(&mut self) {
|
||||
self.current_step += 1;
|
||||
trace!("Warmup scheduler stepped to: {}", self.current_step);
|
||||
}
|
||||
|
||||
fn current_step(&self) -> usize {
|
||||
self.current_step
|
||||
}
|
||||
|
||||
fn reset(&mut self) {
|
||||
self.current_step = 0;
|
||||
debug!("Reset warmup scheduler");
|
||||
}
|
||||
|
||||
fn scheduler_type(&self) -> &'static str {
|
||||
"Warmup"
|
||||
}
|
||||
|
||||
fn base_lr(&self) -> f64 {
|
||||
self.base_lr
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_warmup_creation() {
|
||||
let scheduler = WarmupScheduler::new(0.001, 1000).unwrap();
|
||||
assert_eq!(scheduler.base_lr(), 0.001);
|
||||
assert_eq!(scheduler.warmup_steps(), 1000);
|
||||
assert_eq!(scheduler.current_step(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_warmup_invalid_params() {
|
||||
assert!(WarmupScheduler::new(-0.001, 1000).is_err());
|
||||
assert!(WarmupScheduler::new(0.001, 0).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_warmup_progression() {
|
||||
let scheduler = WarmupScheduler::new(0.001, 100).unwrap();
|
||||
|
||||
// At step 0
|
||||
let lr_0 = scheduler.get_lr(0, 0);
|
||||
assert_eq!(lr_0, 0.0);
|
||||
|
||||
// At step 50 (halfway)
|
||||
let lr_50 = scheduler.get_lr(0, 50);
|
||||
assert!((lr_50 - 0.0005).abs() < 1e-10);
|
||||
|
||||
// At step 100 (end of warmup)
|
||||
let lr_100 = scheduler.get_lr(0, 100);
|
||||
assert!((lr_100 - 0.001).abs() < 1e-10);
|
||||
|
||||
// After warmup
|
||||
let lr_150 = scheduler.get_lr(0, 150);
|
||||
assert!((lr_150 - 0.001).abs() < 1e-10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_warmup_progress() {
|
||||
let scheduler = WarmupScheduler::new(0.001, 100).unwrap();
|
||||
|
||||
assert_eq!(scheduler.warmup_progress(0), 0.0);
|
||||
assert_eq!(scheduler.warmup_progress(50), 0.5);
|
||||
assert_eq!(scheduler.warmup_progress(100), 1.0);
|
||||
assert_eq!(scheduler.warmup_progress(150), 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_warmup_in_warmup() {
|
||||
let scheduler = WarmupScheduler::new(0.001, 100).unwrap();
|
||||
|
||||
assert!(scheduler.in_warmup(0));
|
||||
assert!(scheduler.in_warmup(50));
|
||||
assert!(scheduler.in_warmup(99));
|
||||
assert!(!scheduler.in_warmup(100));
|
||||
assert!(!scheduler.in_warmup(150));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user