use burn_core::data::dataloader::Progress;
use burn_optim::lr_scheduler::module_lr_scheduler::ModuleLearningRate;
use crate::{
LearnerSummary,
renderer::{EvaluationName, MetricsRenderer},
};
pub enum LearnerEvent<T> {
Start {
total_epochs: usize,
starting_epoch: usize,
},
ProcessedItem(TrainingItem<T>),
StartSplit {
epoch_number: usize,
total_items: usize,
},
EndSplit(usize),
EndEpoch(usize),
End(Option<LearnerSummary>),
}
pub enum EvaluatorEvent<T> {
Start {
total_tests: usize,
},
StartTest(EvaluationName, usize),
ProcessedItem(EvaluationName, EvaluationItem<T>),
EndTest,
End(Option<LearnerSummary>),
}
pub trait ItemLazy: Send {
fn sync(self) -> Self;
}
pub trait EventProcessorTraining<TrainEvent, ValidEvent>: Send {
fn process_train(&mut self, event: TrainEvent);
fn process_valid(&mut self, event: ValidEvent);
fn renderer(self) -> Box<dyn MetricsRenderer>;
}
pub trait EventProcessorEvaluation: Send {
type ItemTest: ItemLazy;
fn process_test(&mut self, event: EvaluatorEvent<Self::ItemTest>);
fn renderer(self) -> Box<dyn MetricsRenderer>;
}
#[derive(new)]
pub struct TrainingItem<T> {
pub item: T,
pub progress: Progress,
pub iteration: Option<usize>,
pub lr: Option<ModuleLearningRate>,
}
impl<T: ItemLazy> ItemLazy for TrainingItem<T> {
fn sync(self) -> Self {
TrainingItem {
item: self.item.sync(),
progress: self.progress,
iteration: self.iteration,
lr: self.lr,
}
}
}
#[derive(new)]
pub struct EvaluationItem<T> {
pub item: T,
pub progress: Progress,
pub iteration: Option<usize>,
}
impl<T: ItemLazy> ItemLazy for EvaluationItem<T> {
fn sync(self) -> Self {
EvaluationItem {
item: self.item.sync(),
progress: self.progress,
iteration: self.iteration,
}
}
}
impl ItemLazy for () {
fn sync(self) -> Self {}
}