use crate::{
EarlyStoppingStrategyRef, Interrupter, Learner, LearnerModel, LearnerSummaryConfig,
LearningCheckpointer, LearningResult, SupervisedTrainingEventProcessor, TrainLoader,
ValidLoader,
metric::{
processor::{EventProcessorTraining, LearnerEvent},
store::EventStoreClient,
},
};
use burn_core::prelude::Device;
use burn_core::tensor::distributed::{DistributedConfig, DistributedContext};
use std::sync::Arc;
pub type CustomLearningStrategy<M> = Arc<dyn SupervisedLearningStrategy<M>>;
#[derive(Clone, Copy, Debug)]
pub enum MultiDeviceOptim {
OptimMainDevice,
OptimSharded,
}
pub enum ExecutionStrategy {
SingleDevice(Device),
MultiDevice(Vec<Device>, MultiDeviceOptim),
DistributedDataParallel {
devices: Vec<Device>,
context: DistributedContext,
},
}
impl ExecutionStrategy {
pub fn main_device(&self) -> &Device {
match self {
ExecutionStrategy::SingleDevice(device) => device,
ExecutionStrategy::MultiDevice(devices, _optim) => &devices[0],
ExecutionStrategy::DistributedDataParallel {
devices,
context: _,
} => &devices[0],
}
}
pub fn single(device: Device) -> Self {
Self::SingleDevice(device)
}
pub fn multi(devices: Vec<Device>, optim: MultiDeviceOptim) -> Self {
Self::MultiDevice(devices, optim)
}
}
impl ExecutionStrategy {
pub fn ddp(devices: Vec<Device>, config: DistributedConfig) -> Self {
let context = DistributedContext::init(devices.clone(), config);
Self::DistributedDataParallel { devices, context }
}
}
pub enum TrainingStrategy<M: LearnerModel> {
Default(ExecutionStrategy),
Custom(CustomLearningStrategy<M>),
}
impl<M: LearnerModel> From<ExecutionStrategy> for TrainingStrategy<M> {
fn from(value: ExecutionStrategy) -> Self {
Self::Default(value)
}
}
impl<M: LearnerModel> Default for TrainingStrategy<M> {
fn default() -> Self {
Self::Default(ExecutionStrategy::SingleDevice(Default::default()))
}
}
pub struct TrainingComponents<M: LearnerModel> {
pub num_epochs: usize,
pub checkpoint: Option<usize>,
pub checkpointer: Option<LearningCheckpointer<M>>,
pub grad_accumulation: Option<usize>,
pub interrupter: Interrupter,
pub early_stopping: Option<EarlyStoppingStrategyRef>,
pub event_processor: SupervisedTrainingEventProcessor<M>,
pub event_store: Arc<EventStoreClient>,
pub summary: Option<LearnerSummaryConfig>,
}
pub trait SupervisedLearningStrategy<M: LearnerModel> {
fn train(
&self,
learner: Learner<M>,
dataloader_train: TrainLoader<M>,
dataloader_valid: ValidLoader<M>,
mut training_components: TrainingComponents<M>,
) -> LearningResult<M> {
let starting_epoch = training_components.checkpoint.unwrap_or(0) + 1;
let summary_config = training_components.summary.clone();
training_components
.event_processor
.process_train(LearnerEvent::Start {
total_epochs: training_components.num_epochs,
starting_epoch,
});
let (model, mut event_processor) = self.fit(
training_components,
learner,
dataloader_train,
dataloader_valid,
starting_epoch,
);
let summary = summary_config.and_then(|summary| {
summary
.init()
.map(|summary| summary.with_model(model.to_string()))
.ok()
});
event_processor.process_train(LearnerEvent::End(summary));
let model = model.valid();
let renderer = event_processor.renderer();
LearningResult::<M> { model, renderer }
}
fn fit(
&self,
training_components: TrainingComponents<M>,
learner: Learner<M>,
dataloader_train: TrainLoader<M>,
dataloader_valid: ValidLoader<M>,
starting_epoch: usize,
) -> (M, SupervisedTrainingEventProcessor<M>);
}