incremental_rs/
learning_rate.rs1#[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}