use crate::checkpoint::{
AsyncCheckpointer, Checkpointer, CheckpointingStrategy, ComposedCheckpointingStrategy,
FileCheckpointer, KeepLastNCheckpoints, MetricCheckpointingStrategy,
};
use crate::learner::EarlyStoppingStrategy;
use crate::learner::base::Interrupter;
use crate::logger::{FileMetricLogger, MetricLogger, TrainingProgressLogger};
use crate::metric::processor::{
AsyncProcessorTraining, FullEventProcessorTraining, MetricsTraining,
};
use crate::metric::store::{Aggregate, Direction, EventStoreClient, LogEventStore, Split};
use crate::metric::{Adaptor, LossMetric, Metric, Numeric};
use crate::multi::MultiDeviceLearningStrategy;
use crate::renderer::{MetricsRenderer, default_renderer};
use crate::single::SingleDeviceTrainingStrategy;
use crate::{
ApplicationLoggerInstaller, EarlyStoppingStrategyRef, ExecutionStrategy,
FileApplicationLoggerInstaller, InferenceModelInput, InferenceModelOutput, InferenceStep,
LearnerEvent, LearnerModel, LearnerSummaryConfig, LearningCheckpointer, LearningResult,
TrainStep, TrainingComponents, TrainingModelInput, TrainingModelOutput, TrainingStrategy,
};
use crate::{Learner, SupervisedLearningStrategy};
use burn_core::data::dataloader::DataLoader;
use burn_core::store::ModuleRecord;
use burn_core::tensor::Device;
use burn_optim::OptimizerRecord;
use burn_optim::lr_scheduler::LrSchedulerRecord;
use std::collections::BTreeSet;
use std::path::{Path, PathBuf};
use std::sync::Arc;
pub type TrainLoader<M> = Arc<dyn DataLoader<TrainingModelInput<M>>>;
pub type ValidLoader<M> = Arc<dyn DataLoader<InferenceModelInput<M>>>;
pub type SupervisedTrainingEventProcessor<M> = AsyncProcessorTraining<
LearnerEvent<TrainingModelOutput<M>>,
LearnerEvent<InferenceModelOutput<M>>,
>;
pub struct SupervisedTraining<M: LearnerModel> {
#[allow(clippy::type_complexity)]
checkpointers: Option<(
AsyncCheckpointer<ModuleRecord>,
AsyncCheckpointer<OptimizerRecord>,
AsyncCheckpointer<LrSchedulerRecord>,
)>,
num_epochs: usize,
checkpoint: Option<usize>,
directory: PathBuf,
grad_accumulation: Option<usize>,
grad_checkpointing: bool,
renderer: Option<Box<dyn MetricsRenderer + 'static>>,
metrics: MetricsTraining<TrainingModelOutput<M>, InferenceModelOutput<M>>,
event_store: LogEventStore,
interrupter: Interrupter,
tracing_logger: Option<Box<dyn ApplicationLoggerInstaller>>,
checkpointer_strategy: Box<dyn CheckpointingStrategy>,
early_stopping: Option<EarlyStoppingStrategyRef>,
training_strategy: Option<TrainingStrategy<M>>,
dataloader_train: TrainLoader<M>,
dataloader_valid: ValidLoader<M>,
summary_metrics: BTreeSet<String>,
summary: bool,
progress_logger: Option<Box<dyn TrainingProgressLogger>>,
}
impl<M: LearnerModel> SupervisedTraining<M> {
pub fn new(
directory: impl AsRef<Path>,
dataloader_train: Arc<dyn DataLoader<<M as TrainStep>::Input>>,
dataloader_valid: Arc<dyn DataLoader<<M as InferenceStep>::Input>>,
) -> Self {
let directory = directory.as_ref().to_path_buf();
let experiment_log_file = directory.join("experiment.log");
Self {
num_epochs: 1,
checkpoint: None,
checkpointers: None,
directory,
grad_accumulation: None,
grad_checkpointing: false,
metrics: MetricsTraining::default(),
event_store: LogEventStore::default(),
renderer: None,
interrupter: Interrupter::new(),
tracing_logger: Some(Box::new(FileApplicationLoggerInstaller::new(
experiment_log_file,
))),
checkpointer_strategy: Box::new(
ComposedCheckpointingStrategy::builder()
.add(KeepLastNCheckpoints::new(2))
.add(MetricCheckpointingStrategy::new(
&LossMetric::new(), Aggregate::Mean,
Direction::Lowest,
Split::Valid,
))
.build(),
),
early_stopping: None,
training_strategy: None,
summary_metrics: BTreeSet::new(),
summary: false,
dataloader_train,
dataloader_valid,
progress_logger: None,
}
}
}
impl<M: LearnerModel> SupervisedTraining<M> {
pub fn with_training_strategy(mut self, training_strategy: TrainingStrategy<M>) -> Self {
self.training_strategy = Some(training_strategy);
self
}
pub fn with_metric_logger<ML>(mut self, logger: ML) -> Self
where
ML: MetricLogger + 'static,
{
self.event_store.register_logger(logger);
self
}
pub fn with_progress_logger<PL>(mut self, logger: PL) -> Self
where
PL: TrainingProgressLogger + 'static,
{
self.progress_logger = Some(Box::new(logger));
self
}
pub fn with_checkpointing_strategy<CS: CheckpointingStrategy + 'static>(
mut self,
strategy: CS,
) -> Self {
self.checkpointer_strategy = Box::new(strategy);
self
}
pub fn renderer<MR>(mut self, renderer: MR) -> Self
where
MR: MetricsRenderer + 'static,
{
self.renderer = Some(Box::new(renderer));
self
}
pub fn metrics<Me: MetricRegistration<M>>(self, metrics: Me) -> Self {
metrics.register(self)
}
pub fn metrics_text<Me: TextMetricRegistration<M>>(self, metrics: Me) -> Self {
metrics.register(self)
}
pub fn metric_train<Me: Metric + 'static>(mut self, metric: Me) -> Self
where
TrainingModelOutput<M>: Adaptor<Me::Input>,
{
self.metrics.register_train_metric(metric);
self
}
pub fn metric_valid<Me: Metric + 'static>(mut self, metric: Me) -> Self
where
InferenceModelOutput<M>: Adaptor<Me::Input>,
{
self.metrics.register_valid_metric(metric);
self
}
pub fn grads_accumulation(mut self, accumulation: usize) -> Self {
self.grad_accumulation = Some(accumulation);
self
}
pub fn gradient_checkpointing(mut self) -> Self {
self.grad_checkpointing = true;
self
}
pub fn metric_train_numeric<Me>(mut self, metric: Me) -> Self
where
Me: Metric + Numeric + 'static,
TrainingModelOutput<M>: Adaptor<Me::Input>,
{
self.summary_metrics.insert(metric.name().to_string());
self.metrics.register_train_metric_numeric(metric);
self
}
pub fn metric_valid_numeric<Me: Metric + Numeric + 'static>(mut self, metric: Me) -> Self
where
InferenceModelOutput<M>: Adaptor<Me::Input>,
{
self.summary_metrics.insert(metric.name().to_string());
self.metrics.register_valid_metric_numeric(metric);
self
}
pub fn num_epochs(mut self, num_epochs: usize) -> Self {
self.num_epochs = num_epochs;
self
}
pub fn checkpoint(mut self, checkpoint: usize) -> Self {
self.checkpoint = Some(checkpoint);
self
}
pub fn interrupter(&self) -> Interrupter {
self.interrupter.clone()
}
pub fn with_interrupter(mut self, interrupter: Interrupter) -> Self {
self.interrupter = interrupter;
self
}
pub fn early_stopping<Strategy>(mut self, strategy: Strategy) -> Self
where
Strategy: EarlyStoppingStrategy + Clone + Send + Sync + 'static,
{
self.early_stopping = Some(Box::new(strategy));
self
}
pub fn with_application_logger(
mut self,
logger: Option<Box<dyn ApplicationLoggerInstaller>>,
) -> Self {
self.tracing_logger = logger;
self
}
pub fn with_default_checkpointers(mut self) -> Self {
let checkpoint_dir = self.directory.join("checkpoint");
let checkpointer_model = FileCheckpointer::new(&checkpoint_dir, "model");
let checkpointer_optimizer = FileCheckpointer::new(&checkpoint_dir, "optim");
let checkpointer_scheduler = FileCheckpointer::new(&checkpoint_dir, "scheduler");
self.checkpointers = Some((
AsyncCheckpointer::new(checkpointer_model),
AsyncCheckpointer::new(checkpointer_optimizer),
AsyncCheckpointer::new(checkpointer_scheduler),
));
self
}
pub fn with_custom_checkpointers<CM, CO, CL>(
mut self,
module_checkpointer: CM,
optimizer_checkpointer: CO,
lr_checkpointer: CL,
) -> Self
where
CM: Checkpointer<ModuleRecord> + 'static,
CO: Checkpointer<OptimizerRecord> + 'static,
CL: Checkpointer<LrSchedulerRecord> + 'static,
{
self.checkpointers = Some((
AsyncCheckpointer::new(module_checkpointer),
AsyncCheckpointer::new(optimizer_checkpointer),
AsyncCheckpointer::new(lr_checkpointer),
));
self
}
pub fn summary(mut self) -> Self {
self.summary = true;
self
}
}
impl<M: LearnerModel> SupervisedTraining<M> {
pub fn launch(mut self, learner: Learner<M>) -> LearningResult<M> {
if self.tracing_logger.is_some()
&& let Err(e) = self.tracing_logger.as_ref().unwrap().install()
{
log::warn!("Failed to install the experiment logger: {e}");
}
let renderer = self
.renderer
.unwrap_or_else(|| default_renderer(self.interrupter.clone(), self.checkpoint));
if !self.event_store.has_loggers() {
self.event_store
.register_logger(FileMetricLogger::new(self.directory.clone()));
}
let event_store = Arc::new(EventStoreClient::new(self.event_store));
let full_processor =
FullEventProcessorTraining::new(self.metrics, renderer, event_store.clone());
let full_processor = match self.progress_logger {
Some(logger) => full_processor.with_progress_logger(logger),
None => full_processor,
};
let event_processor = AsyncProcessorTraining::new(full_processor);
let checkpointer = self.checkpointers.map(|(model, optim, scheduler)| {
LearningCheckpointer::new(
model.with_interrupter(self.interrupter.clone()),
optim.with_interrupter(self.interrupter.clone()),
scheduler.with_interrupter(self.interrupter.clone()),
self.checkpointer_strategy,
)
});
let summary = if self.summary {
Some(LearnerSummaryConfig {
directory: self.directory,
metrics: self.summary_metrics.into_iter().collect::<Vec<_>>(),
})
} else {
None
};
let components = TrainingComponents {
checkpoint: self.checkpoint,
checkpointer,
interrupter: self.interrupter,
early_stopping: self.early_stopping,
event_processor,
event_store,
num_epochs: self.num_epochs,
grad_accumulation: self.grad_accumulation,
summary,
};
let training_strategy = self.training_strategy.unwrap_or(TrainingStrategy::Default(
ExecutionStrategy::SingleDevice(autodiff_device(
learner.model.devices()[0].clone(),
self.grad_checkpointing,
)),
));
let mut learner = learner;
if let Some(checkpoint) = components.checkpoint
&& let Some(checkpointer) = &components.checkpointer
{
learner = checkpointer.load_checkpoint(learner, checkpoint);
}
match training_strategy {
TrainingStrategy::Custom(learning_paradigm) => learning_paradigm.train(
learner,
self.dataloader_train,
self.dataloader_valid,
components,
),
TrainingStrategy::Default(strategy) => match strategy {
ExecutionStrategy::SingleDevice(device) => {
let single_device = SingleDeviceTrainingStrategy::new(autodiff_device(
device,
self.grad_checkpointing,
));
single_device.train(
learner,
self.dataloader_train,
self.dataloader_valid,
components,
)
}
ExecutionStrategy::MultiDevice(devices, multi_device_optim) => {
let strategy: Box<dyn SupervisedLearningStrategy<M>> = match devices.len() == 1
{
true => Box::new(SingleDeviceTrainingStrategy::new(autodiff_device(
devices[0].clone(),
self.grad_checkpointing,
))),
false => Box::new(MultiDeviceLearningStrategy::new(
devices
.into_iter()
.map(|d| autodiff_device(d, self.grad_checkpointing))
.collect(),
multi_device_optim,
)),
};
strategy.train(
learner,
self.dataloader_train,
self.dataloader_valid,
components,
)
}
ExecutionStrategy::DistributedDataParallel { devices, context } => {
use crate::ddp::DdpTrainingStrategy;
let ddp = DdpTrainingStrategy::new(
devices
.into_iter()
.map(|d| autodiff_device(d, self.grad_checkpointing))
.collect(),
context,
);
ddp.train(
learner,
self.dataloader_train,
self.dataloader_valid,
components,
)
}
},
}
}
}
fn autodiff_device(mut device: Device, grad_checkpointing: bool) -> Device {
if !device.is_autodiff() {
device = device.autodiff();
}
if grad_checkpointing {
device = device.gradient_checkpointing();
}
device
}
pub trait MetricRegistration<M: LearnerModel>: Sized {
fn register(self, builder: SupervisedTraining<M>) -> SupervisedTraining<M>;
}
pub trait TextMetricRegistration<M: LearnerModel>: Sized {
fn register(self, builder: SupervisedTraining<M>) -> SupervisedTraining<M>;
}
macro_rules! gen_tuple {
($($M:ident),*) => {
impl<$($M,)* M: LearnerModel> TextMetricRegistration<M> for ($($M,)*)
where
$(TrainingModelOutput<M>: Adaptor<$M::Input>,)*
$(InferenceModelOutput<M>: Adaptor<$M::Input>,)*
$($M: Metric + 'static,)*
{
#[allow(non_snake_case)]
fn register(
self,
builder: SupervisedTraining<M>,
) -> SupervisedTraining<M> {
let ($($M,)*) = self;
$(let builder = builder.metric_train($M.clone());)*
$(let builder = builder.metric_valid($M);)*
builder
}
}
impl<$($M,)* M: LearnerModel> MetricRegistration<M> for ($($M,)*)
where
$(TrainingModelOutput<M>: Adaptor<$M::Input>,)*
$(InferenceModelOutput<M>: Adaptor<$M::Input>,)*
$($M: Metric + Numeric + 'static,)*
{
#[allow(non_snake_case)]
fn register(
self,
builder: SupervisedTraining<M>,
) -> SupervisedTraining<M> {
let ($($M,)*) = self;
$(let builder = builder.metric_train_numeric($M.clone());)*
$(let builder = builder.metric_valid_numeric($M);)*
builder
}
}
};
}
gen_tuple!(M1);
gen_tuple!(M1, M2);
gen_tuple!(M1, M2, M3);
gen_tuple!(M1, M2, M3, M4);
gen_tuple!(M1, M2, M3, M4, M5);
gen_tuple!(M1, M2, M3, M4, M5, M6);