use std::collections::HashMap;
use super::{ItemLazy, TrainingItem};
use crate::{
EvaluationItem,
metric::{
Adaptor, Metric, MetricDefinition, MetricEntry, MetricId, MetricMetadata, Numeric,
store::{MetricsUpdate, NumericMetricUpdate},
},
};
pub(crate) struct MetricsTraining<T: ItemLazy, V: ItemLazy> {
train: Vec<Box<dyn MetricUpdater<T>>>,
valid: Vec<Box<dyn MetricUpdater<V>>>,
train_numeric: Vec<Box<dyn NumericMetricUpdater<T>>>,
valid_numeric: Vec<Box<dyn NumericMetricUpdater<V>>>,
metric_definitions: Vec<MetricDefinition>,
}
pub(crate) struct MetricsEvaluation<T: ItemLazy> {
test: Vec<Box<dyn MetricUpdater<T>>>,
test_numeric: Vec<Box<dyn NumericMetricUpdater<T>>>,
metric_definitions: HashMap<MetricId, MetricDefinition>,
}
impl<T: ItemLazy> Default for MetricsEvaluation<T> {
fn default() -> Self {
Self {
test: Default::default(),
test_numeric: Default::default(),
metric_definitions: HashMap::default(),
}
}
}
impl<T: ItemLazy, V: ItemLazy> Default for MetricsTraining<T, V> {
fn default() -> Self {
Self {
train: Vec::default(),
valid: Vec::default(),
train_numeric: Vec::default(),
valid_numeric: Vec::default(),
metric_definitions: Vec::default(),
}
}
}
impl<T: ItemLazy> MetricsEvaluation<T> {
pub(crate) fn register_test_metric<Me: Metric + 'static>(&mut self, metric: Me)
where
T: Adaptor<Me::Input> + 'static,
{
let metric = MetricWrapper::new(metric);
self.register_definition(&metric);
self.test.push(Box::new(metric))
}
pub(crate) fn register_test_metric_numeric<Me: Metric + Numeric + 'static>(
&mut self,
metric: Me,
) where
T: Adaptor<Me::Input> + 'static,
{
let metric = MetricWrapper::new(metric);
self.register_definition(&metric);
self.test_numeric.push(Box::new(metric))
}
fn register_definition<Me: Metric>(&mut self, metric: &MetricWrapper<Me>) {
self.metric_definitions.insert(
metric.id.clone(),
MetricDefinition::new(metric.id.clone(), &metric.metric),
);
}
pub(crate) fn metric_definitions(&mut self) -> Vec<MetricDefinition> {
self.metric_definitions.values().cloned().collect()
}
pub(crate) fn update_test(
&mut self,
item: &EvaluationItem<T>,
metadata: &MetricMetadata,
) -> MetricsUpdate {
let mut entries = Vec::with_capacity(self.test.len());
let mut entries_numeric = Vec::with_capacity(self.test_numeric.len());
for metric in self.test.iter_mut() {
let state = metric.update(&item.item, metadata);
entries.push(state);
}
for metric in self.test_numeric.iter_mut() {
let numeric_update = metric.update(&item.item, metadata);
entries_numeric.push(numeric_update);
}
MetricsUpdate::new(entries, entries_numeric)
}
}
impl<T: ItemLazy, V: ItemLazy> MetricsTraining<T, V> {
pub(crate) fn register_train_metric<Me: Metric + 'static>(&mut self, metric: Me)
where
T: Adaptor<Me::Input> + 'static,
{
let metric = MetricWrapper::new(metric);
self.register_definition(&metric);
self.train.push(Box::new(metric))
}
pub(crate) fn register_valid_metric<Me: Metric + 'static>(&mut self, metric: Me)
where
V: Adaptor<Me::Input> + 'static,
{
let metric = MetricWrapper::new(metric);
self.register_definition(&metric);
self.valid.push(Box::new(metric))
}
pub(crate) fn register_train_metric_numeric<Me: Metric + Numeric + 'static>(
&mut self,
metric: Me,
) where
T: Adaptor<Me::Input> + 'static,
{
let metric = MetricWrapper::new(metric);
self.register_definition(&metric);
self.train_numeric.push(Box::new(metric))
}
pub(crate) fn register_valid_metric_numeric<Me>(&mut self, metric: Me)
where
V: Adaptor<Me::Input> + 'static,
Me: Metric + Numeric + 'static,
{
let metric = MetricWrapper::new(metric);
self.register_definition(&metric);
self.valid_numeric.push(Box::new(metric))
}
fn register_definition<Me: Metric>(&mut self, metric: &MetricWrapper<Me>) {
if !self
.metric_definitions
.iter()
.any(|def| def.metric_id == metric.id)
{
self.metric_definitions
.push(MetricDefinition::new(metric.id.clone(), &metric.metric));
}
}
pub(crate) fn metric_definitions(&mut self) -> Vec<MetricDefinition> {
self.metric_definitions.clone()
}
pub(crate) fn update_train(
&mut self,
item: &TrainingItem<T>,
metadata: &MetricMetadata,
) -> MetricsUpdate {
let mut entries = Vec::with_capacity(self.train.len());
let mut entries_numeric = Vec::with_capacity(self.train_numeric.len());
for metric in self.train.iter_mut() {
let state = metric.update(&item.item, metadata);
entries.push(state);
}
for metric in self.train_numeric.iter_mut() {
let numeric_update = metric.update(&item.item, metadata);
entries_numeric.push(numeric_update);
}
MetricsUpdate::new(entries, entries_numeric)
}
pub(crate) fn update_valid(
&mut self,
item: &TrainingItem<V>,
metadata: &MetricMetadata,
) -> MetricsUpdate {
let mut entries = Vec::with_capacity(self.valid.len());
let mut entries_numeric = Vec::with_capacity(self.valid_numeric.len());
for metric in self.valid.iter_mut() {
let state = metric.update(&item.item, metadata);
entries.push(state);
}
for metric in self.valid_numeric.iter_mut() {
let numeric_update = metric.update(&item.item, metadata);
entries_numeric.push(numeric_update);
}
MetricsUpdate::new(entries, entries_numeric)
}
pub(crate) fn end_epoch_train(&mut self) -> MetricsUpdate {
let mut entries = Vec::with_capacity(self.train.len());
let mut entries_numeric = Vec::with_capacity(self.train_numeric.len());
for metric in self.train.iter_mut() {
entries.push(metric.compute());
metric.clear();
}
for metric in self.train_numeric.iter_mut() {
entries_numeric.push(metric.compute());
metric.clear();
}
MetricsUpdate::new(entries, entries_numeric)
}
pub(crate) fn end_epoch_valid(&mut self) -> MetricsUpdate {
let mut entries = Vec::with_capacity(self.valid.len());
let mut entries_numeric = Vec::with_capacity(self.valid_numeric.len());
for metric in self.valid.iter_mut() {
entries.push(metric.compute());
metric.clear();
}
for metric in self.valid_numeric.iter_mut() {
entries_numeric.push(metric.compute());
metric.clear();
}
MetricsUpdate::new(entries, entries_numeric)
}
}
impl<T> From<&TrainingItem<T>> for MetricMetadata {
fn from(item: &TrainingItem<T>) -> Self {
Self {
progress: item.progress.clone(),
iteration: item.iteration,
lr: item.lr.clone(),
}
}
}
impl<T> From<&EvaluationItem<T>> for MetricMetadata {
fn from(item: &EvaluationItem<T>) -> Self {
Self {
progress: item.progress.clone(),
iteration: item.iteration,
lr: None,
}
}
}
pub(crate) trait NumericMetricUpdater<T>: Send + Sync {
fn update(&mut self, item: &T, metadata: &MetricMetadata) -> NumericMetricUpdate;
fn compute(&mut self) -> NumericMetricUpdate;
fn clear(&mut self);
}
pub(crate) trait MetricUpdater<T>: Send + Sync {
fn update(&mut self, item: &T, metadata: &MetricMetadata) -> MetricEntry;
fn compute(&mut self) -> MetricEntry;
fn clear(&mut self);
}
pub(crate) struct MetricWrapper<M> {
pub id: MetricId,
pub metric: M,
}
impl<M: Metric> MetricWrapper<M> {
pub fn new(metric: M) -> Self {
Self {
id: MetricId::new(metric.name()),
metric,
}
}
}
impl<T, M> NumericMetricUpdater<T> for MetricWrapper<M>
where
T: 'static,
M: Metric + Numeric + 'static,
T: Adaptor<M::Input>,
{
fn update(&mut self, item: &T, metadata: &MetricMetadata) -> NumericMetricUpdate {
let serialized_entry = self.metric.update(&item.adapt(), metadata);
let update = MetricEntry::new(self.id.clone(), serialized_entry);
let numeric = self.metric.value();
let running = self.metric.running_value();
NumericMetricUpdate {
entry: update,
numeric_entry: numeric,
running_entry: running,
}
}
fn compute(&mut self) -> NumericMetricUpdate {
let serialized_entry = self.metric.compute();
let update = MetricEntry::new(self.id.clone(), serialized_entry);
let final_entry = self.metric.final_value();
NumericMetricUpdate {
entry: update,
numeric_entry: Some(final_entry),
running_entry: None,
}
}
fn clear(&mut self) {
self.metric.clear()
}
}
impl<T, M> MetricUpdater<T> for MetricWrapper<M>
where
T: 'static,
M: Metric + 'static,
T: Adaptor<M::Input>,
{
fn update(&mut self, item: &T, metadata: &MetricMetadata) -> MetricEntry {
let serialized_entry = self.metric.update(&item.adapt(), metadata);
MetricEntry::new(self.id.clone(), serialized_entry)
}
fn compute(&mut self) -> MetricEntry {
let serialized_entry = self.metric.compute();
MetricEntry::new(self.id.clone(), serialized_entry)
}
fn clear(&mut self) {
self.metric.clear()
}
}