use std::fmt::Debug;
use alloc::collections::BTreeMap;
pub(super) use alloc::string::String;
use alloc::vec::Vec;
use burn_core as burn;
use burn_core::config::Config;
use crate::lr_scheduler::composed::ComposedLrSchedulerConfig;
use crate::lr_scheduler::cosine::CosineAnnealingLrSchedulerConfig;
use crate::lr_scheduler::exponential::ExponentialLrSchedulerConfig;
use crate::lr_scheduler::linear::LinearLrSchedulerConfig;
use crate::lr_scheduler::noam::NoamLrSchedulerConfig;
use crate::lr_scheduler::sequential::SequentialLrSchedulerConfig;
use crate::lr_scheduler::step::StepLrSchedulerConfig;
use crate::{RecordState, StateSink, StateSource, join_path};
use burn::store::RecordError;
use burn::tensor::{Bytes, Device};
use burn_pack::{Reader, Scalar, Writer};
use crate::LearningRate;
macro_rules! impl_from_for_scheduler {
($($variant:ident($config:ident)),* $(,)?) => {
$(
impl From<$config> for LrSchedulerConfig {
fn from(config: $config) -> Self {
LrSchedulerConfig::$variant(config)
}
}
)*
};
}
pub trait LrScheduler: LrSchedulerClone + Send + Sync {
fn step(&mut self) -> LearningRate;
fn to_record(&self) -> LrSchedulerRecord;
fn load_record(&mut self, record: LrSchedulerRecord);
}
pub trait LrSchedulerClone {
fn clone_box(&self) -> Box<dyn LrScheduler>;
}
impl<T> LrSchedulerClone for T
where
T: 'static + LrScheduler + Clone,
{
fn clone_box(&self) -> Box<dyn LrScheduler> {
Box::new(self.clone())
}
}
impl Clone for Box<dyn LrScheduler> {
fn clone(&self) -> Box<dyn LrScheduler> {
self.as_ref().clone_box()
}
}
#[derive(Clone)]
pub struct DynLrScheduler {
scheduler: Box<dyn LrScheduler>,
}
impl DynLrScheduler {
pub fn step(&mut self) -> LearningRate {
self.scheduler.step()
}
pub fn to_record(&self) -> LrSchedulerRecord {
self.scheduler.to_record()
}
pub fn load_record(mut self, record: LrSchedulerRecord) -> Self {
self.scheduler.load_record(record);
self
}
}
impl<S> From<S> for DynLrScheduler
where
S: LrScheduler + 'static,
{
fn from(scheduler: S) -> Self {
Self {
scheduler: Box::new(scheduler),
}
}
}
#[derive(Default, Clone, Debug)]
pub struct LrSchedulerRecord {
scalars: BTreeMap<String, Scalar>,
}
impl LrSchedulerRecord {
pub fn new() -> Self {
Self::default()
}
pub fn is_empty(&self) -> bool {
self.scalars.is_empty()
}
pub fn with_scalar<V: Into<Scalar>>(mut self, key: &str, value: V) -> Self {
self.scalars.insert(String::from(key), value.into());
self
}
pub fn scalar<V: TryFrom<Scalar>>(&self, key: &str) -> Option<V> {
self.scalars
.get(key)
.copied()
.and_then(|scalar| V::try_from(scalar).ok())
}
pub fn with_record(mut self, prefix: &str, record: LrSchedulerRecord) -> Self {
for (key, value) in record.scalars {
self.scalars.insert(join_path(prefix, &key), value);
}
self
}
pub fn record(&self, prefix: &str) -> LrSchedulerRecord {
let head = join_path(prefix, "");
let scalars = self
.scalars
.iter()
.filter_map(|(key, value)| {
key.strip_prefix(&head)
.map(|stripped| (String::from(stripped), *value))
})
.collect();
LrSchedulerRecord { scalars }
}
pub fn from_state<S: RecordState>(state: &S) -> Self {
let mut sink = StateSink::default();
state.state_flatten("", &mut sink);
debug_assert!(
sink.tensors.is_empty(),
"learning rate scheduler state is expected to be scalar-only"
);
Self {
scalars: sink.scalars.into_iter().collect(),
}
}
pub fn into_state<S: RecordState>(&self) -> Option<S> {
let mut source = StateSource::new(self.scalars.clone());
S::state_unflatten("", &mut source, &Device::default())
}
pub fn into_bytes(self) -> Result<Bytes, RecordError> {
Ok(self.into_writer().into_bytes()?)
}
pub fn from_bytes(bytes: Bytes) -> Result<Self, RecordError> {
let reader = Reader::from_bytes(bytes)?;
Ok(Self {
scalars: reader.scalars().clone(),
})
}
#[cfg(feature = "std")]
pub fn save<P: AsRef<std::path::Path>>(self, path: P) -> Result<(), RecordError> {
self.into_writer().write_to_file(path)?;
Ok(())
}
#[cfg(feature = "std")]
pub fn load<P: AsRef<std::path::Path>>(path: P) -> Result<Self, RecordError> {
let reader = Reader::from_file(path)?;
Ok(Self {
scalars: reader.scalars().clone(),
})
}
fn into_writer(self) -> Writer {
let mut writer = Writer::new(Vec::new());
for (key, value) in &self.scalars {
writer = writer.with_scalar(key, *value);
}
writer
}
}
#[derive(Config, Debug)]
pub enum LrSchedulerConfig {
Constant(LearningRate),
Linear(LinearLrSchedulerConfig),
Cosine(CosineAnnealingLrSchedulerConfig),
Exponential(ExponentialLrSchedulerConfig),
Noam(NoamLrSchedulerConfig),
Step(StepLrSchedulerConfig),
Composed(ComposedLrSchedulerConfig),
Sequential(SequentialLrSchedulerConfig),
}
impl LrSchedulerConfig {
pub(crate) fn build(&self) -> Result<DynLrScheduler, String> {
Ok(match self {
Self::Constant(lr) => (*lr).into(),
Self::Linear(config) => config.build()?.into(),
Self::Cosine(config) => config.build()?.into(),
Self::Exponential(config) => config.build()?.into(),
Self::Noam(config) => config.build()?.into(),
Self::Step(config) => config.build()?.into(),
Self::Composed(config) => config.build()?.into(),
Self::Sequential(config) => config.build()?.into(),
})
}
}
impl_from_for_scheduler!(
Constant(LearningRate),
Linear(LinearLrSchedulerConfig),
Cosine(CosineAnnealingLrSchedulerConfig),
Exponential(ExponentialLrSchedulerConfig),
Noam(NoamLrSchedulerConfig),
Step(StepLrSchedulerConfig),
Composed(ComposedLrSchedulerConfig),
Sequential(SequentialLrSchedulerConfig),
);
#[cfg(test)]
pub(super) mod test_utils {
use super::*;
const LOOSE_EPSILON: LearningRate = 1e-10;
pub fn check_lr_sequence<I, S>(mut scheduler: S, expected_lrs: I)
where
I: IntoIterator<Item = LearningRate>,
S: LrScheduler,
{
expected_lrs
.into_iter()
.enumerate()
.for_each(|(i, expected)| {
let lr = scheduler.step();
assert!(
(lr - expected).abs() < LOOSE_EPSILON,
"Scheduled learning rate {lr} is not approximately equal to the expected value \
{expected} at step {i}",
);
});
}
pub fn check_save_load<S>(mut scheduler: S, save_at_step: usize)
where
S: Clone + LrScheduler,
{
let mut truth = scheduler.clone();
(0..save_at_step).for_each(|_| {
truth.step();
scheduler.step();
});
let rec = scheduler.to_record();
scheduler.load_record(rec);
compare_steps(&mut scheduler, &mut truth, save_at_step);
}
pub fn compare_steps<S: LrScheduler>(a: &mut S, b: &mut S, num_steps: usize) {
(0..num_steps).for_each(|i| {
let lr_a = a.step();
let lr_b = b.step();
assert!(
(lr_a - lr_b).abs() < LOOSE_EPSILON,
"The two learning rates ({lr_a}, {lr_b}) at position {i} in the remaining \
sequences are not approximately equal",
);
});
}
}