Skip to main content

incremental_rs/
learning_rate.rs

1#[derive(Debug, Clone)]
2pub enum LearningRateSchedule {
3    Constant {
4        initial_rate: f64,
5    },
6    InverseScaling {
7        initial_rate: f64,
8        decay: f64,
9        power: f64,
10    },
11}
12
13impl LearningRateSchedule {
14    pub fn calculate(&self, step: usize) -> f64 {
15        match *self {
16            LearningRateSchedule::Constant { initial_rate } => initial_rate,
17            LearningRateSchedule::InverseScaling {
18                initial_rate,
19                decay,
20                power,
21            } => initial_rate / (1.0 + decay * (step as f64)).powf(power),
22        }
23    }
24}