use std::sync::Arc;
use crate::{
LearnerSummary,
logger::{EvaluationProgressLogger, TrainingProgressLogger},
metric::{MetricDefinition, MetricEntry, NumericEntry},
};
pub trait MetricsRendererTraining: Send + Sync {
fn update_train(&mut self, state: MetricState);
fn update_valid(&mut self, state: MetricState);
fn on_train_end(
&mut self,
summary: Option<LearnerSummary>,
) -> Result<(), Box<dyn core::error::Error>> {
default_summary_action(summary);
Ok(())
}
}
pub trait MetricsRenderer:
MetricsRendererEvaluation
+ MetricsRendererTraining
+ TrainingProgressLogger
+ EvaluationProgressLogger
{
fn manual_close(&mut self);
fn register_metric(&mut self, definition: MetricDefinition);
}
#[derive(Clone)]
pub struct EvaluationName {
pub(crate) name: Arc<String>,
}
impl EvaluationName {
pub fn new<S: core::fmt::Display>(s: S) -> Self {
Self {
name: Arc::new(format!("{s}")),
}
}
pub fn as_str(&self) -> &str {
&self.name
}
}
impl core::fmt::Display for EvaluationName {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(&self.name)
}
}
pub trait MetricsRendererEvaluation: Send + Sync {
fn update_test(&mut self, name: EvaluationName, state: MetricState);
fn on_test_end(
&mut self,
summary: Option<LearnerSummary>,
) -> Result<(), Box<dyn core::error::Error>> {
default_summary_action(summary);
Ok(())
}
}
#[derive(Debug)]
pub enum MetricState {
Generic(MetricEntry),
Numeric(MetricEntry, Option<NumericEntry>),
}
fn default_summary_action(summary: Option<LearnerSummary>) {
if let Some(summary) = summary {
println!("{summary}");
}
}