use burn_core::{self as burn};
use super::cosine::CosineAnnealingLrSchedulerConfig;
use super::exponential::ExponentialLrSchedulerConfig;
use super::linear::LinearLrSchedulerConfig;
use super::noam::NoamLrSchedulerConfig;
use super::{LrScheduler, LrSchedulerRecord, String};
use crate::LearningRate;
use crate::lr_scheduler::module_lr_scheduler::ModuleLrScheduler;
use crate::lr_scheduler::step::StepLrSchedulerConfig;
use crate::lr_scheduler::{DynLrScheduler, LrSchedulerConfig};
use burn::config::Config;
#[derive(Config, Debug)]
pub struct ComposedLrSchedulerConfig {
#[config(default = "Vec::new()")]
schedulers: Vec<LrSchedulerConfig>,
#[config(default = "SchedulerReduction::Prod")]
reduction: SchedulerReduction,
}
#[derive(Clone)]
pub struct ComposedLrScheduler {
schedulers: Vec<DynLrScheduler>,
reduction: SchedulerReduction,
}
#[derive(Config, Debug, Copy)]
pub enum SchedulerReduction {
Avg,
Sum,
Prod,
}
impl ComposedLrSchedulerConfig {
pub(crate) fn build(&self) -> Result<ComposedLrScheduler, String> {
let mut schedulers: Vec<DynLrScheduler> = Vec::with_capacity(self.schedulers.len());
for config in self.schedulers.iter() {
schedulers.push(config.build()?);
}
Ok(ComposedLrScheduler {
schedulers,
reduction: self.reduction,
})
}
pub fn init(&self) -> Result<ModuleLrScheduler, String> {
self.build().map(|s| s.into())
}
pub fn constant(mut self, lr: LearningRate) -> Self {
self.schedulers.push(LrSchedulerConfig::Constant(lr));
self
}
pub fn linear(mut self, config: LinearLrSchedulerConfig) -> Self {
self.schedulers.push(LrSchedulerConfig::Linear(config));
self
}
pub fn cosine(mut self, config: CosineAnnealingLrSchedulerConfig) -> Self {
self.schedulers.push(LrSchedulerConfig::Cosine(config));
self
}
pub fn exponential(mut self, config: ExponentialLrSchedulerConfig) -> Self {
self.schedulers.push(LrSchedulerConfig::Exponential(config));
self
}
pub fn noam(mut self, config: NoamLrSchedulerConfig) -> Self {
self.schedulers.push(LrSchedulerConfig::Noam(config));
self
}
pub fn step(mut self, config: StepLrSchedulerConfig) -> Self {
self.schedulers.push(LrSchedulerConfig::Step(config));
self
}
pub fn composed(mut self, config: Self) -> Self {
self.schedulers.push(LrSchedulerConfig::Composed(config));
self
}
}
impl ComposedLrScheduler {
pub fn with_custom_scheduler<S: LrScheduler + 'static>(mut self, scheduler: S) -> Self {
self.schedulers.push(scheduler.into());
self
}
}
impl LrScheduler for ComposedLrScheduler {
fn step(&mut self) -> LearningRate {
let mut step = match self.reduction {
SchedulerReduction::Avg => 0.0,
SchedulerReduction::Sum => 0.0,
SchedulerReduction::Prod => 1.0,
};
let num_scheduler = self.schedulers.len() as f64;
for lr in self.schedulers.iter_mut().map(|s| s.step()) {
step = match self.reduction {
SchedulerReduction::Avg => step + (lr / num_scheduler),
SchedulerReduction::Sum => step + lr,
SchedulerReduction::Prod => step * lr,
}
}
step
}
fn to_record(&self) -> LrSchedulerRecord {
let mut record = LrSchedulerRecord::new();
for (index, item) in self.schedulers.iter().enumerate() {
let sub = item.to_record();
record = record.with_record(&index.to_string(), sub);
}
record
}
fn load_record(&mut self, record: LrSchedulerRecord) {
self.schedulers = self
.schedulers
.clone()
.iter_mut()
.enumerate()
.map(|(index, item)| {
let sub = record.record(&index.to_string());
item.clone().load_record(sub)
})
.collect();
}
}