use crate::{
Learner, LearnerEvent, LearnerModel, MultiDeviceOptim, SupervisedLearningStrategy,
SupervisedTrainingEventProcessor, TrainLoader, TrainingComponents, ValidLoader,
metric::processor::EventProcessorTraining,
multi::epoch::MultiDeviceTrainEpoch,
single::{TrainingLoop, epoch::SingleDeviceValidEpoch},
};
use burn_core::{data::dataloader::split::split_dataloader, tensor::Device};
pub struct MultiDeviceLearningStrategy {
devices: Vec<Device>,
optim: MultiDeviceOptim,
}
impl MultiDeviceLearningStrategy {
pub fn new(devices: Vec<Device>, optim: MultiDeviceOptim) -> Self {
Self { devices, optim }
}
}
impl<M: LearnerModel> SupervisedLearningStrategy<M> for MultiDeviceLearningStrategy {
fn fit(
&self,
training_components: TrainingComponents<M>,
mut learner: Learner<M>,
dataloader_train: TrainLoader<M>,
dataloader_valid: ValidLoader<M>,
starting_epoch: usize,
) -> (M, SupervisedTrainingEventProcessor<M>) {
let main_device = self.devices.first().unwrap();
let train_total_items = dataloader_train.num_items();
let dataloader_train = split_dataloader(dataloader_train, &self.devices);
let dataloader_valid = dataloader_valid.to_device(&main_device.clone().inner());
let valid_total_items = dataloader_valid.num_items();
learner.fork(main_device);
let mut event_processor = training_components.event_processor;
let mut checkpointer = training_components.checkpointer;
let mut early_stopping = training_components.early_stopping;
let epoch_train = MultiDeviceTrainEpoch::<M>::new(
dataloader_train.clone(),
training_components.grad_accumulation,
);
let epoch_valid: SingleDeviceValidEpoch<M> =
SingleDeviceValidEpoch::new(dataloader_valid.clone());
for training_progress in TrainingLoop::new(starting_epoch, training_components.num_epochs) {
let epoch = training_progress.items_processed;
event_processor.process_train(LearnerEvent::StartSplit {
epoch_number: epoch,
total_items: train_total_items,
});
epoch_train.run(
&mut learner,
&training_progress,
&mut event_processor,
&training_components.interrupter,
self.devices.to_vec(),
self.optim,
);
event_processor.process_train(LearnerEvent::EndSplit(epoch));
if training_components.interrupter.should_stop() {
let reason = training_components
.interrupter
.get_message()
.unwrap_or(String::from("Reason unknown"));
log::info!("Training interrupted: {reason}");
break;
}
if matches!(self.optim, MultiDeviceOptim::OptimSharded) {
learner.fork(main_device);
}
event_processor.process_valid(LearnerEvent::StartSplit {
epoch_number: epoch,
total_items: valid_total_items,
});
epoch_valid.run(
&learner,
&training_progress,
&mut event_processor,
&training_components.interrupter,
);
event_processor.process_valid(LearnerEvent::EndSplit(epoch));
event_processor.process_train(LearnerEvent::EndEpoch(epoch));
if let Some(checkpointer) = &mut checkpointer {
checkpointer.checkpoint(&learner, epoch, &training_components.event_store);
}
if let Some(early_stopping) = &mut early_stopping
&& early_stopping.should_stop(epoch, &training_components.event_store)
{
break;
}
}
(learner.model(), event_processor)
}
}