use super::{EventProcessorTraining, ItemLazy, LearnerEvent, MetricsTraining};
use crate::{
logger::TrainingProgressLogger,
metric::store::{EpochSummary, EventStoreClient, Split},
renderer::cli::CliMetricsRenderer,
};
use std::sync::Arc;
#[allow(dead_code)]
pub(crate) struct MinimalEventProcessor<T: ItemLazy, V: ItemLazy> {
metrics: MetricsTraining<T, V>,
store: Arc<EventStoreClient>,
progress_logger: Option<Box<dyn TrainingProgressLogger>>,
}
#[allow(dead_code)]
impl<T: ItemLazy, V: ItemLazy> MinimalEventProcessor<T, V> {
pub(crate) fn new(metrics: MetricsTraining<T, V>, store: Arc<EventStoreClient>) -> Self {
Self {
metrics,
store,
progress_logger: None,
}
}
pub(crate) fn with_progress_logger(mut self, logger: Box<dyn TrainingProgressLogger>) -> Self {
self.progress_logger = Some(logger);
self
}
}
impl<T: ItemLazy, V: ItemLazy> EventProcessorTraining<LearnerEvent<T>, LearnerEvent<V>>
for MinimalEventProcessor<T, V>
{
fn process_train(&mut self, event: LearnerEvent<T>) {
match event {
LearnerEvent::Start {
total_epochs,
starting_epoch,
} => {
let definitions = self.metrics.metric_definitions();
self.store
.add_event_train(crate::metric::store::Event::MetricsInit(definitions));
if let Some(logger) = &mut self.progress_logger {
logger.start(total_epochs, starting_epoch, None);
}
}
LearnerEvent::StartSplit {
epoch_number,
total_items,
} => {
self.store
.add_event_train(crate::metric::store::Event::StartSplit(epoch_number));
if let Some(logger) = &mut self.progress_logger {
logger.start_split(Split::Train.into(), total_items);
}
}
LearnerEvent::ProcessedItem(item) => {
let item = item.sync();
let metadata = (&item).into();
let update = self.metrics.update_train(&item, &metadata);
self.store
.add_event_train(crate::metric::store::Event::MetricsUpdate(update));
if let Some(logger) = &mut self.progress_logger {
logger.update_split(item.progress.items_processed);
}
}
LearnerEvent::EndSplit(epoch) => {
let update = self.metrics.end_epoch_train();
self.store
.add_event_train(crate::metric::store::Event::MetricsUpdate(update));
self.store
.add_event_train(crate::metric::store::Event::EndEpoch(EpochSummary::new(
epoch,
Split::Train,
)));
if let Some(logger) = &mut self.progress_logger {
logger.end_split();
}
}
LearnerEvent::EndEpoch(epoch) => {
if let Some(logger) = &mut self.progress_logger {
logger.update_epoch(epoch);
}
}
LearnerEvent::End(_summary) => {
if let Some(logger) = &mut self.progress_logger {
logger.end();
}
}
}
}
fn process_valid(&mut self, event: LearnerEvent<V>) {
match event {
LearnerEvent::Start { .. } => {} LearnerEvent::StartSplit {
epoch_number,
total_items,
} => {
self.store
.add_event_valid(crate::metric::store::Event::StartSplit(epoch_number));
if let Some(logger) = &mut self.progress_logger {
logger.start_split(Split::Valid.into(), total_items);
}
}
LearnerEvent::ProcessedItem(item) => {
let item = item.sync();
let metadata = (&item).into();
let update = self.metrics.update_valid(&item, &metadata);
self.store
.add_event_valid(crate::metric::store::Event::MetricsUpdate(update));
if let Some(logger) = &mut self.progress_logger {
logger.update_split(item.progress.items_processed);
}
}
LearnerEvent::EndSplit(epoch) => {
let update = self.metrics.end_epoch_valid();
self.store
.add_event_valid(crate::metric::store::Event::MetricsUpdate(update));
self.store
.add_event_valid(crate::metric::store::Event::EndEpoch(EpochSummary::new(
epoch,
Split::Valid,
)));
if let Some(logger) = &mut self.progress_logger {
logger.end_split();
}
}
LearnerEvent::EndEpoch(_) => {} LearnerEvent::End(_) => {} }
}
fn renderer(self) -> Box<dyn crate::renderer::MetricsRenderer> {
Box::new(CliMetricsRenderer::new())
}
}