burn-optim 0.22.0-pre.1

Optimizer building blocks for the Burn deep learning framework
Documentation
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;

/// Compose multiple [learning rate schedulers](LrScheduler) together.
#[derive(Config, Debug)]
pub struct ComposedLrSchedulerConfig {
    #[config(default = "Vec::new()")]
    schedulers: Vec<LrSchedulerConfig>,
    #[config(default = "SchedulerReduction::Prod")]
    reduction: SchedulerReduction,
}

/// Compose multiple [learning rate schedulers](LrScheduler) together.
#[derive(Clone)]
pub struct ComposedLrScheduler {
    schedulers: Vec<DynLrScheduler>,
    reduction: SchedulerReduction,
}

/// Defines how the learning rates generated by the schedulers are combined.
#[derive(Config, Debug, Copy)]
pub enum SchedulerReduction {
    /// All learning rates are averaged.
    Avg,
    /// All learning rates are summed.
    Sum,
    /// All learning rates are multiplied.
    Prod,
}

impl ComposedLrSchedulerConfig {
    /// Initialize the learning rate scheduler.
    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,
        })
    }

    /// Initializes a [module learning rate scheduler](ModuleLrScheduler).
    pub fn init(&self) -> Result<ModuleLrScheduler, String> {
        self.build().map(|s| s.into())
    }

    /// Appends a [constant learning rate](crate::lr_scheduler::constant::ConstantLr).
    pub fn constant(mut self, lr: LearningRate) -> Self {
        self.schedulers.push(LrSchedulerConfig::Constant(lr));
        self
    }

    /// Appends a [linear scheduler](crate::lr_scheduler::linear::LinearLrScheduler).
    pub fn linear(mut self, config: LinearLrSchedulerConfig) -> Self {
        self.schedulers.push(LrSchedulerConfig::Linear(config));
        self
    }

    /// Appends a [cosine scheduler](ComposedLrSchedulerConfig).
    pub fn cosine(mut self, config: CosineAnnealingLrSchedulerConfig) -> Self {
        self.schedulers.push(LrSchedulerConfig::Cosine(config));
        self
    }

    /// Appends an [exponential scheduler](crate::lr_scheduler::exponential::ExponentialLrScheduler).
    pub fn exponential(mut self, config: ExponentialLrSchedulerConfig) -> Self {
        self.schedulers.push(LrSchedulerConfig::Exponential(config));
        self
    }

    /// Appends a [noam scheduler](crate::lr_scheduler::noam::NoamLrScheduler).
    pub fn noam(mut self, config: NoamLrSchedulerConfig) -> Self {
        self.schedulers.push(LrSchedulerConfig::Noam(config));
        self
    }

    /// Appends a [step scheduler](crate::lr_scheduler::step::StepLrScheduler).
    pub fn step(mut self, config: StepLrSchedulerConfig) -> Self {
        self.schedulers.push(LrSchedulerConfig::Step(config));
        self
    }

    /// Appends a [composed scheduler](ComposedLrScheduler).
    pub fn composed(mut self, config: Self) -> Self {
        self.schedulers.push(LrSchedulerConfig::Composed(config));
        self
    }
}

impl ComposedLrScheduler {
    /// Add a custom learning rate scheduler to existing ones.
    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();
    }
}