diff --git a/pywr-core/src/agg_funcs/mod.rs b/pywr-core/src/agg_funcs/mod.rs index e69ac567..c77449b1 100644 --- a/pywr-core/src/agg_funcs/mod.rs +++ b/pywr-core/src/agg_funcs/mod.rs @@ -4,7 +4,7 @@ mod py; #[cfg(feature = "pyo3")] pub use py::PyAggFunc; -use crate::recorders::PeriodValue; +use crate::recorders::{Event, PeriodValue}; use thiserror::Error; #[derive(Error, Debug)] @@ -85,6 +85,85 @@ impl AggFuncF64 { } } + /// Calculate the aggregation over the given slice of `Event`. + /// + /// This function computes the aggregation based on the duration of each event in fraction days. + /// Only completed events (those with a defined end time) are included. + /// It returns an `Option`, which will be `None` if the aggregation cannot be computed (e.g., for `Mean` with no events). + pub fn calc_events(&self, events: &[Event]) -> Option { + match self { + Self::Sum => Some( + events + .iter() + .filter_map(|e| e.duration().map(|d| d.fractional_days())) + .sum(), + ), + Self::Mean => { + let total_duration: f64 = events + .iter() + .filter_map(|e| e.duration().map(|d| d.fractional_days())) + .sum(); + let count = events.len() as f64; + if count == 0.0 { + None + } else { + Some(total_duration / count) + } + } + Self::Min => events + .iter() + .filter_map(|e| e.duration().map(|d| d.fractional_days())) + .min_by(|a, b| { + a.partial_cmp(b) + .expect("Failed to calculate minimum of event durations containing a NaN.") + }), + Self::Max => events + .iter() + .filter_map(|e| e.duration().map(|d| d.fractional_days())) + .max_by(|a, b| { + a.partial_cmp(b) + .expect("Failed to calculate maximum of event durations containing a NaN.") + }), + Self::CountNonZero => { + let count = events.iter().filter(|e| e.end.is_some()).count(); + Some(count as f64) + } + Self::CountFunc { func } => { + let count = events + .iter() + .filter(|e| e.duration().map(|d| func(d.fractional_days())).unwrap_or(false)) + .count(); + Some(count as f64) + } + Self::Product => { + let product = events + .iter() + .filter_map(|e| e.duration().map(|d| d.fractional_days())) + .product(); + Some(product) + } + Self::AnyNonZero { tolerance } => { + let any = events.iter().any(|e| { + e.duration() + .map(|d| d.fractional_days().abs() > *tolerance) + .unwrap_or(false) + }); + Some(any as u8 as f64) + } + #[cfg(feature = "pyo3")] + Self::Python(py_func) => { + let vals: Vec = events + .iter() + .filter_map(|e| e.duration().map(|d| d.fractional_days())) + .collect(); + match py_func.call_f64(vals) { + Ok(result) => Some(result), + Err(e) => panic!("Error in Python aggregation function: {}", e), + } + } + } + } + /// Calculate the aggregation of the given iterator of values. pub fn calc_iter_f64<'a, V>(&self, values: V) -> Result where diff --git a/pywr-core/src/lib.rs b/pywr-core/src/lib.rs index b7dcde61..7cb4ad9e 100644 --- a/pywr-core/src/lib.rs +++ b/pywr-core/src/lib.rs @@ -18,6 +18,7 @@ pub mod models; pub mod network; pub mod node; pub mod parameters; +pub mod predicate; pub mod recorders; pub mod scenario; pub mod solvers; diff --git a/pywr-core/src/parameters/mod.rs b/pywr-core/src/parameters/mod.rs index 0a1c366f..c3999cf3 100644 --- a/pywr-core/src/parameters/mod.rs +++ b/pywr-core/src/parameters/mod.rs @@ -81,7 +81,7 @@ use std::hash::{Hash, Hasher}; use std::marker::PhantomData; use std::ops::Deref; use thiserror::Error; -pub use threshold::{Predicate, ThresholdParameter}; +pub use threshold::ThresholdParameter; pub use vector::VectorParameter; /// Simple parameter index. diff --git a/pywr-core/src/parameters/multi_threshold.rs b/pywr-core/src/parameters/multi_threshold.rs index ada18cf0..b7778620 100644 --- a/pywr-core/src/parameters/multi_threshold.rs +++ b/pywr-core/src/parameters/multi_threshold.rs @@ -2,8 +2,9 @@ use crate::metric::MetricF64; use crate::network::Network; use crate::parameters::errors::{ParameterCalculationError, ParameterSetupError}; use crate::parameters::{ - GeneralParameter, Parameter, ParameterMeta, ParameterName, ParameterState, Predicate, downcast_internal_state_mut, + GeneralParameter, Parameter, ParameterMeta, ParameterName, ParameterState, downcast_internal_state_mut, }; +use crate::predicate::Predicate; use crate::scenario::ScenarioIndex; use crate::state::State; use crate::timestep::Timestep; @@ -100,7 +101,8 @@ impl GeneralParameter for MultiThresholdParameter { mod tests { use super::MultiThresholdParameter; use crate::metric::MetricF64; - use crate::parameters::{Array1Parameter, Predicate}; + use crate::parameters::Array1Parameter; + use crate::predicate::Predicate; use crate::test_utils::{run_and_assert_parameter_u64, simple_model}; use ndarray::{Array1, Array2, Axis, concatenate}; diff --git a/pywr-core/src/parameters/threshold.rs b/pywr-core/src/parameters/threshold.rs index 777aff8d..daff252e 100644 --- a/pywr-core/src/parameters/threshold.rs +++ b/pywr-core/src/parameters/threshold.rs @@ -4,31 +4,11 @@ use crate::parameters::errors::{ParameterCalculationError, ParameterSetupError}; use crate::parameters::{ GeneralParameter, Parameter, ParameterMeta, ParameterName, ParameterState, downcast_internal_state_mut, }; +use crate::predicate::Predicate; use crate::scenario::ScenarioIndex; use crate::state::State; use crate::timestep::Timestep; -pub enum Predicate { - LessThan, - GreaterThan, - EqualTo, - LessThanOrEqualTo, - GreaterThanOrEqualTo, -} - -impl Predicate { - /// Apply the predicate to a value and a threshold. - pub fn apply(&self, value: f64, threshold: f64) -> bool { - match self { - Predicate::LessThan => value < threshold, - Predicate::GreaterThan => value > threshold, - Predicate::EqualTo => (value - threshold).abs() < 1E-6, // TODO make this a global constant - Predicate::LessThanOrEqualTo => value <= threshold, - Predicate::GreaterThanOrEqualTo => value >= threshold, - } - } -} - pub struct ThresholdParameter { meta: ParameterMeta, metric: MetricF64, diff --git a/pywr-core/src/predicate.rs b/pywr-core/src/predicate.rs new file mode 100644 index 00000000..36a91231 --- /dev/null +++ b/pywr-core/src/predicate.rs @@ -0,0 +1,20 @@ +#[derive(Debug, Clone)] +pub enum Predicate { + LessThan, + GreaterThan, + EqualTo, + LessThanOrEqualTo, + GreaterThanOrEqualTo, +} + +impl Predicate { + pub fn apply(&self, a: f64, b: f64) -> bool { + match self { + Predicate::LessThan => a < b, + Predicate::GreaterThan => a > b, + Predicate::EqualTo => (a - b).abs() < 1E-6, // TODO make this a global constant + Predicate::LessThanOrEqualTo => a <= b, + Predicate::GreaterThanOrEqualTo => a >= b, + } + } +} diff --git a/pywr-core/src/recorders/aggregator/agg_func.rs b/pywr-core/src/recorders/aggregator/agg_func.rs new file mode 100644 index 00000000..28ef8e03 --- /dev/null +++ b/pywr-core/src/recorders/aggregator/agg_func.rs @@ -0,0 +1,87 @@ +use crate::recorders::aggregator::{Event, PeriodValue}; + +#[derive(Clone, Debug)] +pub enum AggregationFunction { + Sum, + Mean, + Min, + Max, + CountNonZero, + CountFunc { func: fn(f64) -> bool }, +} + +impl AggregationFunction { + /// Calculate the aggregation of the given `PeriodValue`. + /// + /// This function takes a slice of `PeriodValue` and applies the aggregation function to the values. + /// It returns an `Option`, which will be `None` if the aggregation cannot be computed (e.g., for `Mean` with no values). + /// + pub fn calc_period_values(&self, values: &[PeriodValue]) -> Option { + match self { + AggregationFunction::Sum => Some(values.iter().map(|v| v.value * v.duration.fractional_days()).sum()), + AggregationFunction::Mean => { + let ndays: f64 = values.iter().map(|v| v.duration.fractional_days()).sum(); + if ndays == 0.0 { + None + } else { + let sum: f64 = values.iter().map(|v| v.value * v.duration.fractional_days()).sum(); + + Some(sum / ndays) + } + } + AggregationFunction::Min => values.iter().map(|v| v.value).min_by(|a, b| { + a.partial_cmp(b) + .expect("Failed to calculate minimum of values containing a NaN.") + }), + AggregationFunction::Max => values.iter().map(|v| v.value).max_by(|a, b| { + a.partial_cmp(b) + .expect("Failed to calculate maximum of values containing a NaN.") + }), + AggregationFunction::CountNonZero => { + let count = values.iter().filter(|v| v.value != 0.0).count(); + Some(count as f64) + } + AggregationFunction::CountFunc { func } => { + let count = values.iter().filter(|v| func(v.value)).count(); + Some(count as f64) + } + } + } + + pub fn calc_f64(&self, values: &[f64]) -> Option { + match self { + AggregationFunction::Sum => Some(values.iter().sum()), + AggregationFunction::Mean => { + let ndays: i64 = values.len() as i64; + if ndays == 0 { + None + } else { + let sum: f64 = values.iter().sum(); + Some(sum / ndays as f64) + } + } + AggregationFunction::Min => values + .iter() + .min_by(|a, b| { + a.partial_cmp(b) + .expect("Failed to calculate minimum of values containing a NaN.") + }) + .copied(), + AggregationFunction::Max => values + .iter() + .max_by(|a, b| { + a.partial_cmp(b) + .expect("Failed to calculate maximum of values containing a NaN.") + }) + .copied(), + AggregationFunction::CountNonZero => { + let count = values.iter().filter(|v| **v != 0.0).count(); + Some(count as f64) + } + AggregationFunction::CountFunc { func } => { + let count = values.iter().filter(|v| func(**v)).count(); + Some(count as f64) + } + } + } +} diff --git a/pywr-core/src/recorders/aggregator/event.rs b/pywr-core/src/recorders/aggregator/event.rs new file mode 100644 index 00000000..c9c078c7 --- /dev/null +++ b/pywr-core/src/recorders/aggregator/event.rs @@ -0,0 +1,245 @@ +use crate::predicate::Predicate; +use crate::recorders::aggregator::PeriodValue; +use crate::timestep::PywrDuration; +use chrono::NaiveDateTime; + +#[derive(Default, Clone, Debug)] +enum EventState { + #[default] + Ended, + Started(NaiveDateTime), +} + +#[derive(Debug, Clone, Copy)] +pub struct Event { + pub start: NaiveDateTime, + pub end: Option, +} + +impl Event { + pub fn duration(&self) -> Option { + self.end.map(|end| (end - self.start).into()) + } +} + +#[derive(Default, Debug, Clone)] +pub struct EventAggregatorState { + current: EventState, +} + +#[derive(Debug, Clone)] +pub struct EventAggregator { + predicate: Predicate, + threshold: f64, +} + +impl EventAggregator { + pub fn new(predicate: Predicate, threshold: f64) -> Self { + Self { predicate, threshold } + } + + pub fn setup(&self) -> EventAggregatorState { + EventAggregatorState::default() + } + + /// Process a new value and return an event if one has completed. + pub fn process_value(&self, current_state: &mut EventAggregatorState, value: PeriodValue) -> Option { + let active_now = self.predicate.apply(value.value, self.threshold); + + let (new_current, event) = match (¤t_state.current, active_now) { + (EventState::Ended, true) => { + // Start a new event + (EventState::Started(value.start), None) + } + (EventState::Started(started), false) => { + // End the current event + let event = Event { + start: *started, + end: Some(value.start), + }; + + (EventState::Ended, Some(event)) + } + (EventState::Started(started), true) => { + // Continue the current event + (EventState::Started(*started), None) + } + (EventState::Ended, false) => { + // No event to continue + (EventState::Ended, None) + } + }; + + current_state.current = new_current; + + event + } +} + +#[cfg(test)] +mod tests { + use super::{EventAggregator, EventAggregatorState}; + use crate::predicate::Predicate; + use crate::recorders::aggregator::PeriodValue; + use chrono::{NaiveDate, TimeDelta}; + + #[test] + fn test_event_aggregator() { + let agg = EventAggregator { + predicate: Predicate::GreaterThan, + threshold: 1.0, + }; + + let mut state = EventAggregatorState::default(); + + let start = NaiveDate::from_ymd_opt(2023, 1, 30) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 3.0)); + assert!(agg_value.is_none()); + + let start = NaiveDate::from_ymd_opt(2023, 1, 31) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 3.0)); + assert!(agg_value.is_none()); + + let start = NaiveDate::from_ymd_opt(2023, 2, 1) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); + assert!(agg_value.is_some()); + + let start = NaiveDate::from_ymd_opt(2023, 2, 2) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); + assert!(agg_value.is_none()); + } + + #[test] + fn test_event_aggregator_less_than() { + let agg = EventAggregator { + predicate: Predicate::LessThan, + threshold: 2.0, + }; + let mut state = EventAggregatorState::default(); + + let start = NaiveDate::from_ymd_opt(2023, 3, 1) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let v1 = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.5)); + assert!(v1.is_none()); + + let start2 = NaiveDate::from_ymd_opt(2023, 3, 2) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let v2 = agg.process_value(&mut state, PeriodValue::new(start2, TimeDelta::days(1).into(), 2.5)); + assert!(v2.is_some()); + let event = v2.unwrap(); + assert_eq!(event.start, start); + assert_eq!(event.end, Some(start2)); + } + + #[test] + fn test_multiple_events() { + let agg = EventAggregator { + predicate: Predicate::GreaterThan, + threshold: 5.0, + }; + let mut state = EventAggregatorState::default(); + + let dates: Vec<_> = (0..6) + .map(|i| { + NaiveDate::from_ymd_opt(2023, 4, 1 + i) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap() + }) + .collect(); + + // Start event + assert!( + agg.process_value(&mut state, PeriodValue::new(dates[0], TimeDelta::days(1).into(), 6.0)) + .is_none() + ); + // Continue event + assert!( + agg.process_value(&mut state, PeriodValue::new(dates[1], TimeDelta::days(1).into(), 7.0)) + .is_none() + ); + // End event + let ev1 = agg.process_value(&mut state, PeriodValue::new(dates[2], TimeDelta::days(1).into(), 4.0)); + assert!(ev1.is_some()); + assert_eq!(ev1.unwrap().start, dates[0]); + assert_eq!(ev1.unwrap().end, Some(dates[2])); + + // Start new event + assert!( + agg.process_value(&mut state, PeriodValue::new(dates[3], TimeDelta::days(1).into(), 8.0)) + .is_none() + ); + // End new event + let ev2 = agg.process_value(&mut state, PeriodValue::new(dates[4], TimeDelta::days(1).into(), 2.0)); + assert!(ev2.is_some()); + assert_eq!(ev2.unwrap().start, dates[3]); + assert_eq!(ev2.unwrap().end, Some(dates[4])); + + // No event + assert!( + agg.process_value(&mut state, PeriodValue::new(dates[5], TimeDelta::days(1).into(), 1.0)) + .is_none() + ); + } + + #[test] + fn test_no_event_triggered() { + let agg = EventAggregator { + predicate: Predicate::GreaterThan, + threshold: 10.0, + }; + let mut state = EventAggregatorState::default(); + + for i in 0..5 { + let start = NaiveDate::from_ymd_opt(2023, 5, 1 + i) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let v = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 5.0)); + assert!(v.is_none()); + } + } + + #[test] + fn test_event_starts_but_never_ends() { + let agg = EventAggregator { + predicate: Predicate::GreaterThan, + threshold: 2.0, + }; + let mut state = EventAggregatorState::default(); + + let start = NaiveDate::from_ymd_opt(2023, 6, 1) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + assert!( + agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 3.0)) + .is_none() + ); + + let start2 = NaiveDate::from_ymd_opt(2023, 6, 2) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + assert!( + agg.process_value(&mut state, PeriodValue::new(start2, TimeDelta::days(1).into(), 4.0)) + .is_none() + ); + } +} diff --git a/pywr-core/src/recorders/aggregator/mod.rs b/pywr-core/src/recorders/aggregator/mod.rs index 7c3c2a59..a96be6a6 100644 --- a/pywr-core/src/recorders/aggregator/mod.rs +++ b/pywr-core/src/recorders/aggregator/mod.rs @@ -1,302 +1,159 @@ -use crate::agg_funcs::AggFuncF64; -use crate::timestep::PywrDuration; -use chrono::{Datelike, Duration, NaiveDate, NaiveDateTime, NaiveTime}; -use std::num::NonZeroUsize; +mod event; +mod periodic; -#[derive(Clone, Debug)] -pub enum AggregationFrequency { - Monthly, - Annual, - Days(NonZeroUsize), +use crate::recorders::metric_set::MetricSetOutputInfo; +use crate::timestep::TimeDomain; +pub use event::{Event, EventAggregator, EventAggregatorState}; +use periodic::PeriodicAggregatorState; +pub use periodic::{AggregationFrequency, PeriodValue, PeriodicAggregator}; + +#[derive(Debug, Clone)] +pub enum AggregatorState { + Periodic(PeriodicAggregatorState), + Event(EventAggregatorState), } -impl AggregationFrequency { - fn is_date_in_period(&self, period_start: &NaiveDateTime, date: &NaiveDateTime) -> bool { +impl AggregatorState { + fn as_periodic(&self) -> Option<&PeriodicAggregatorState> { match self { - Self::Monthly => (period_start.year() == date.year()) && (period_start.month() == date.month()), - Self::Annual => period_start.year() == date.year(), - Self::Days(days) => { - let period_end = *period_start + Duration::days(days.get() as i64); - (period_start <= date) && (date < &period_end) - } + AggregatorState::Periodic(state) => Some(state), + _ => None, } } - fn start_of_next_period(&self, current_date: &NaiveDateTime) -> NaiveDateTime { + fn as_periodic_mut(&mut self) -> Option<&mut PeriodicAggregatorState> { match self { - Self::Monthly => { - let current_month = current_date.month(); - // Increment the year if we're in December - let year = if current_month == 12 { - current_date.year() + 1 - } else { - current_date.year() - }; - let next_month = (current_month % 12) + 1; - // 1st of the next month - // SAFETY: This should be safe to unwrap as it will always create a valid date unless - // we are at the limit of dates that are representable. - let date = NaiveDate::from_ymd_opt(year, next_month, 1).unwrap(); - NaiveDateTime::new(date, NaiveTime::default()) - } - Self::Annual => { - // 1st of January in the next year - // SAFETY: This should be safe to unwrap as it will always create a valid date unless - // we are at the limit of dates that are representable. - let date = NaiveDate::from_ymd_opt(current_date.year() + 1, 1, 1).unwrap(); - NaiveDateTime::new(date, NaiveTime::default()) - } - Self::Days(days) => *current_date + Duration::days(days.get() as i64), + AggregatorState::Periodic(state) => Some(state), + _ => None, } } - /// Split the value representing a period into multiple ['PeriodValue'] that do not cross the - /// boundary of the given period. - fn split_value_into_periods(&self, value: PeriodValue) -> Vec> { - let mut sub_values = Vec::new(); - - let mut current_date = value.start; - let end_date = value.duration + value.start; - - while current_date < end_date { - let start_of_next_period = self.start_of_next_period(¤t_date); - - let current_duration = if start_of_next_period <= end_date { - start_of_next_period - current_date - } else { - end_date - current_date - }; - - sub_values.push(PeriodValue { - start: current_date, - duration: current_duration.into(), - value: value.value, - }); - - current_date = start_of_next_period; + fn as_event_mut(&mut self) -> Option<&mut EventAggregatorState> { + match self { + AggregatorState::Event(state) => Some(state), + _ => None, } - - sub_values } } -#[derive(Default, Debug, Clone)] -struct PeriodicAggregatorState { - current_values: Option>>, +#[derive(Debug, Clone)] +pub struct NestedAggregatorState { + state: AggregatorState, + child: Option>, } -impl PeriodicAggregatorState { - fn process_value( - &mut self, - value: PeriodValue, - agg_freq: &AggregationFrequency, - agg_func: &AggFuncF64, - ) -> Option> { - if let Some(current_values) = self.current_values.as_mut() { - // SAFETY: The current_values vector is guaranteed to contain at least one value. - let current_period_start = current_values - .first() - .expect("Aggregation state contains no values when at least one is expected.") - .start; - - // Determine if the value is in the current period - if agg_freq.is_date_in_period(¤t_period_start, &value.start) { - // New value in the current aggregation period; just append it. - current_values.push(value); - - None - } else { - // New value is part of a different period (assume the next one). - - // Calculate the aggregated value of the previous period. - let agg_period = if let Some(agg_value) = agg_func.calc_period_values(current_values) { - let agg_duration = value.start - current_period_start; - Some(PeriodValue::new(current_period_start, agg_duration.into(), agg_value)) - } else { - None - }; - - // Reset the state for the next period - current_values.clear(); - current_values.push(value); - - // Finally return the aggregated value from the previous period - agg_period - } - } else { - // No previous values defined; just append the value - self.current_values = Some(vec![value]); - - None - } - } - - fn process_value_no_period(&mut self, value: PeriodValue) { - if let Some(current_values) = self.current_values.as_mut() { - current_values.push(value); - } else { - self.current_values = Some(vec![value]); - } - } +#[derive(Debug, Clone)] +pub enum AggregatorValue { + Periodic(PeriodValue), + Event(Event), +} - fn calc_aggregation(&self, agg_func: &AggFuncF64) -> Option> { - if let Some(current_values) = &self.current_values { - if let Some(agg_value) = agg_func.calc_period_values(current_values) { - // SAFETY: The current_values vector is guaranteed to contain at least one value. - let current_period_start = current_values - .first() - .expect("Aggregation state contains no values when at least one is expected.") - .start; - - let current_period_end = current_values - .last() - .expect("Aggregation state contains no values when at least one is expected.") - .start; - let current_period_duration = current_period_end - current_period_start; - Some(PeriodValue::new( - current_period_start, - current_period_duration.into(), - agg_value, - )) - } else { - None - } - } else { - None - } +impl From for AggregatorValue { + fn from(event: Event) -> Self { + AggregatorValue::Event(event) } } -#[derive(Clone, Debug)] -struct PeriodicAggregator { - frequency: Option, - function: AggFuncF64, +impl From> for AggregatorValue { + fn from(value: PeriodValue) -> Self { + AggregatorValue::Periodic(value) + } } -#[derive(Debug, Copy, Clone)] -pub struct PeriodValue { - pub start: NaiveDateTime, - pub duration: PywrDuration, - pub value: T, +#[derive(Debug, Clone)] +pub enum Aggregator { + Periodic(PeriodicAggregator), + Event(EventAggregator), } -impl PeriodValue { - pub fn new(start: NaiveDateTime, duration: PywrDuration, value: T) -> Self { - Self { start, duration, value } - } - - /// The end of the period. - pub fn end(&self) -> NaiveDateTime { - self.duration + self.start +impl From for Aggregator { + fn from(agg: PeriodicAggregator) -> Self { + Aggregator::Periodic(agg) } } -impl PeriodValue> { - pub fn index(&self, index: usize) -> PeriodValue - where - T: Copy, - { - PeriodValue { - start: self.start, - duration: self.duration, - value: self.value[index], - } - } - pub fn len(&self) -> usize { - self.value.len() +impl From for Aggregator { + fn from(agg: EventAggregator) -> Self { + Aggregator::Event(agg) } +} - pub fn is_empty(&self) -> bool { - self.value.is_empty() +impl From for AggregatorState { + fn from(state: PeriodicAggregatorState) -> Self { + AggregatorState::Periodic(state) } } -impl From<&[PeriodValue]> for PeriodValue> -where - T: Copy, -{ - fn from(values: &[PeriodValue]) -> Self { - let start = values.first().expect("Empty vector of period values.").start; - let duration = values.last().expect("Empty vector of period values.").duration; - - let value = values.iter().map(|v| v.value).collect(); - Self { start, duration, value } +impl From for AggregatorState { + fn from(state: EventAggregatorState) -> Self { + AggregatorState::Event(state) } } -impl PeriodicAggregator { - fn setup(&self) -> PeriodicAggregatorState { - PeriodicAggregatorState::default() +impl Aggregator { + fn setup(&self) -> AggregatorState { + match self { + Aggregator::Periodic(_) => PeriodicAggregatorState::default().into(), + Aggregator::Event(_) => EventAggregatorState::default().into(), + } } - /// Append a new value to the aggregator. - /// - /// The new value should sequentially follow from the previously processed values. If the - /// value completes a new aggregation period then a value representing that aggregation is - /// returned. - fn process_value( - &self, - current_state: &mut PeriodicAggregatorState, - value: PeriodValue, - ) -> Option> { - // Split the given period into separate periods that align with the aggregation period. - let mut agg_value = None; - - if let Some(period) = &self.frequency { - for v in period.split_value_into_periods(value) { - let av = current_state.process_value(v, period, &self.function); - if av.is_some() { - if agg_value.is_some() { - panic!( - "Multiple aggregated values yielded from aggregator. This indicates that the given value spans multiple aggregation periods which is not supported." - ) - } - agg_value = av; - } - } - } else { - current_state.process_value_no_period(value); + fn process_value(&self, state: &mut AggregatorState, value: PeriodValue) -> Option { + match self { + Aggregator::Periodic(agg) => agg + .process_value(state.as_periodic_mut().unwrap(), value) + .map(|v| v.into()), + Aggregator::Event(agg) => agg + .process_value(state.as_event_mut().unwrap(), value) + .map(|v| v.into()), } - agg_value } - fn calc_aggregation(&self, state: &PeriodicAggregatorState) -> Option> { - state.calc_aggregation(&self.function) + fn calc_aggregation(&self, state: &AggregatorState) -> Option { + match self { + Aggregator::Periodic(agg) => agg.calc_aggregation(state.as_periodic().unwrap()).map(|v| v.into()), + Aggregator::Event(_) => None, + } } -} -#[derive(Debug, Clone)] -pub struct AggregatorState { - state: PeriodicAggregatorState, - child: Option>, + fn output_info(&self, time_domain: &TimeDomain) -> MetricSetOutputInfo { + match self { + Aggregator::Periodic(agg) => MetricSetOutputInfo::Periodic { + num_periods: agg.number_of_periods(time_domain), + }, + Aggregator::Event(_) => MetricSetOutputInfo::Event, + } + } } #[derive(Clone, Debug)] -pub struct Aggregator { - agg: PeriodicAggregator, - child: Option>, +pub struct NestedAggregator { + aggregator: Aggregator, + child: Option>, } -impl Aggregator { - pub fn new(period: Option, function: AggFuncF64, child: Option) -> Self { +impl NestedAggregator { + pub fn new(aggregator: Aggregator, child: Option) -> Self { Self { - agg: PeriodicAggregator { - frequency: period, - function, - }, + aggregator, child: child.map(Box::new), } } - pub fn setup(&self) -> AggregatorState { - AggregatorState { - state: self.agg.setup(), + pub fn output_info(&self, time_domain: &TimeDomain) -> MetricSetOutputInfo { + self.aggregator.output_info(time_domain) + } + + /// Create the initial default state for the aggregator. + pub fn setup(&self) -> NestedAggregatorState { + NestedAggregatorState { + state: self.aggregator.setup(), child: self.child.as_ref().map(|c| Box::new(c.setup())), } } /// Append a new value to the aggregator. - pub fn append_value(&self, state: &mut AggregatorState, value: PeriodValue) -> Option> { + pub fn append_value(&self, state: &mut NestedAggregatorState, value: AggregatorValue) -> Option { let agg_value = match (&self.child, state.child.as_mut()) { (Some(child), Some(child_state)) => child.append_value(child_state, value), (None, None) => Some(value), @@ -305,7 +162,14 @@ impl Aggregator { }; if let Some(agg_value) = agg_value { - self.agg.process_value(&mut state.state, agg_value) + match agg_value { + AggregatorValue::Periodic(value) => self.aggregator.process_value(&mut state.state, value), + AggregatorValue::Event(_event) => { + panic!( + "It is not possible to process an event value in a nested aggregator. The event aggregator should be the top level aggregator." + ) + } + } } else { None } @@ -315,7 +179,7 @@ impl Aggregator { /// /// This will also compute the final aggregation value from the child aggregators if any exists. /// This includes aggregation calculations over partial or unfinished periods. - pub fn finalise(&self, state: &mut AggregatorState) -> Option> { + pub fn finalise(&self, state: &mut NestedAggregatorState) -> Option { let final_child_value = match (&self.child, state.child.as_mut()) { (Some(child), Some(child_state)) => child.finalise(child_state), (None, None) => None, @@ -324,89 +188,48 @@ impl Aggregator { }; // If there is a final value from the child aggregator then process it - if let Some(final_child_value) = final_child_value { - let _ = self.agg.process_value(&mut state.state, final_child_value); + if let Some(agg_value) = final_child_value { + match agg_value { + AggregatorValue::Periodic(value) => { + let _ = self.aggregator.process_value(&mut state.state, value); + } + AggregatorValue::Event(_event) => { + panic!( + "It is not possible to process an event value in a nested aggregator. The event aggregator should be the top level aggregator." + ) + } + } } // Finally, compute the aggregation of the current state - self.agg.calc_aggregation(&state.state) - } - - /// Create the initial default state for the aggregator. - pub fn default_state(&self) -> AggregatorState { - let state = PeriodicAggregatorState::default(); - let child = self.child.as_ref().map(|c| Box::new(c.default_state())); - AggregatorState { state, child } + self.aggregator.calc_aggregation(&state.state) } } #[cfg(test)] mod tests { - use super::{AggFuncF64, AggregationFrequency, Aggregator, PeriodicAggregator, PeriodicAggregatorState}; + use super::{AggregationFrequency, Aggregator, AggregatorValue, NestedAggregator, PeriodicAggregator}; + use crate::agg_funcs::AggFuncF64; use crate::recorders::aggregator::PeriodValue; use chrono::{Datelike, NaiveDate, TimeDelta}; use float_cmp::assert_approx_eq; - #[test] - fn test_periodic_aggregator() { - let agg = PeriodicAggregator { - frequency: Some(AggregationFrequency::Monthly), - function: AggFuncF64::Sum, - }; - - let mut state = PeriodicAggregatorState::default(); - - let start = NaiveDate::from_ymd_opt(2023, 1, 30) - .unwrap() - .and_hms_opt(0, 0, 0) - .unwrap(); - let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); - assert!(agg_value.is_none()); - - let start = NaiveDate::from_ymd_opt(2023, 1, 31) - .unwrap() - .and_hms_opt(0, 0, 0) - .unwrap(); - let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); - assert!(agg_value.is_none()); - - let start = NaiveDate::from_ymd_opt(2023, 2, 1) - .unwrap() - .and_hms_opt(0, 0, 0) - .unwrap(); - let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); - assert!(agg_value.is_some()); - - let start = NaiveDate::from_ymd_opt(2023, 2, 2) - .unwrap() - .and_hms_opt(0, 0, 0) - .unwrap(); - let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); - assert!(agg_value.is_none()); - } - #[test] fn test_nested_aggregator() { - let model_agg = PeriodicAggregator { - frequency: None, - function: AggFuncF64::Max, - }; + let model_agg = PeriodicAggregator::new(None, AggFuncF64::Max); - let annual_agg = PeriodicAggregator { - frequency: Some(AggregationFrequency::Annual), - function: AggFuncF64::Min, - }; + let annual_agg = PeriodicAggregator::new(Some(AggregationFrequency::Annual), AggFuncF64::Min); // Setup an aggregator to calculate the max of the annual minimum values - let max_annual_min = Aggregator { - agg: model_agg, - child: Some(Box::new(Aggregator { - agg: annual_agg, + let max_annual_min = NestedAggregator { + aggregator: Aggregator::Periodic(model_agg), + child: Some(Box::new(NestedAggregator { + aggregator: Aggregator::Periodic(annual_agg), child: None, })), }; - let mut state = max_annual_min.default_state(); + let mut state = max_annual_min.setup(); let mut date = NaiveDate::from_ymd_opt(2023, 1, 1) .unwrap() @@ -414,53 +237,19 @@ mod tests { .unwrap(); for _i in 0..365 * 3 { let value = PeriodValue::new(date, TimeDelta::days(1).into(), date.year() as f64); - let _agg_value = max_annual_min.append_value(&mut state, value); + let _agg_value = max_annual_min.append_value(&mut state, value.into()); date += TimeDelta::days(1); } let final_value = max_annual_min.finalise(&mut state); if let Some(final_value) = final_value { - assert_approx_eq!(f64, final_value.value, 2025.0); + match final_value { + AggregatorValue::Periodic(value) => assert_approx_eq!(f64, value.value, 2025.0), + _ => panic!("Final value is not a PeriodValue!"), + } } else { panic!("Final value is None!") } } - - #[test] - fn test_sub_daily_aggregation() { - let values = vec![ - PeriodValue::new( - NaiveDate::from_ymd_opt(2023, 1, 1) - .unwrap() - .and_hms_opt(0, 0, 0) - .unwrap(), - TimeDelta::hours(1).into(), - 2.0, - ), - PeriodValue::new( - NaiveDate::from_ymd_opt(2023, 1, 1) - .unwrap() - .and_hms_opt(1, 0, 0) - .unwrap(), - TimeDelta::hours(2).into(), - 1.0, - ), - PeriodValue::new( - NaiveDate::from_ymd_opt(2023, 1, 1) - .unwrap() - .and_hms_opt(3, 0, 0) - .unwrap(), - TimeDelta::hours(1).into(), - 3.0, - ), - ]; - - let agg_value = AggFuncF64::Mean.calc_period_values(values.as_slice()).unwrap(); - assert_approx_eq!(f64, agg_value, 7.0 / 4.0); - - let agg_value = AggFuncF64::Sum.calc_period_values(values.as_slice()).unwrap(); - let expected = 2.0 + 1.0 + 3.0; - assert_approx_eq!(f64, agg_value, expected); - } } diff --git a/pywr-core/src/recorders/aggregator/periodic.rs b/pywr-core/src/recorders/aggregator/periodic.rs new file mode 100644 index 00000000..d4d748ca --- /dev/null +++ b/pywr-core/src/recorders/aggregator/periodic.rs @@ -0,0 +1,382 @@ +use crate::agg_funcs::AggFuncF64; +use crate::timestep::{PywrDuration, TimeDomain}; +use chrono::{Datelike, Duration, NaiveDate, NaiveDateTime, NaiveTime}; +use std::num::NonZeroUsize; + +#[derive(Clone, Debug)] +pub enum AggregationFrequency { + Monthly, + Annual, + Days(NonZeroUsize), +} + +impl AggregationFrequency { + /// Number of periods in the given time domain. + fn number_of_periods(&self, time_domain: &TimeDomain) -> usize { + let start = time_domain.first().expect("No time-steps in time domain.").date; + let end = time_domain.last().expect("No time-steps in time domain.").date; + + match self { + Self::Monthly => { + let n_years = (end.year() - start.year()) as u32; + (n_years * 12 + end.month() - start.month()) as usize + } + Self::Annual => (end.year() - start.year()) as usize, + Self::Days(days) => { + let n_days = end.signed_duration_since(start).num_days(); + (n_days / days.get() as i64) as usize + } + } + } + + fn is_date_in_period(&self, period_start: &NaiveDateTime, date: &NaiveDateTime) -> bool { + match self { + Self::Monthly => (period_start.year() == date.year()) && (period_start.month() == date.month()), + Self::Annual => period_start.year() == date.year(), + Self::Days(days) => { + let period_end = *period_start + Duration::days(days.get() as i64); + (period_start <= date) && (date < &period_end) + } + } + } + + fn start_of_next_period(&self, current_date: &NaiveDateTime) -> NaiveDateTime { + match self { + Self::Monthly => { + let current_month = current_date.month(); + // Increment the year if we're in December + let year = if current_month == 12 { + current_date.year() + 1 + } else { + current_date.year() + }; + let next_month = (current_month % 12) + 1; + // 1st of the next month + // SAFETY: This should be safe to unwrap as it will always create a valid date unless + // we are at the limit of dates that are representable. + let date = NaiveDate::from_ymd_opt(year, next_month, 1).unwrap(); + NaiveDateTime::new(date, NaiveTime::default()) + } + Self::Annual => { + // 1st of January in the next year + // SAFETY: This should be safe to unwrap as it will always create a valid date unless + // we are at the limit of dates that are representable. + let date = NaiveDate::from_ymd_opt(current_date.year() + 1, 1, 1).unwrap(); + NaiveDateTime::new(date, NaiveTime::default()) + } + Self::Days(days) => *current_date + Duration::days(days.get() as i64), + } + } + + /// Split the value representing a period into multiple ['PeriodValue'] that do not cross the + /// boundary of the given period. + fn split_value_into_periods(&self, value: PeriodValue) -> Vec> { + let mut sub_values = Vec::new(); + + let mut current_date = value.start; + let end_date = value.duration + value.start; + + while current_date < end_date { + let start_of_next_period = self.start_of_next_period(¤t_date); + + let current_duration = if start_of_next_period <= end_date { + start_of_next_period - current_date + } else { + end_date - current_date + }; + + sub_values.push(PeriodValue { + start: current_date, + duration: current_duration.into(), + value: value.value, + }); + + current_date = start_of_next_period; + } + + sub_values + } +} + +/// State of the periodic aggregator. +/// +/// This state stores the current values, if any, that are yielded from the aggregation on the +/// given time-step. Periodic output is consistent for each metric, and therefore is stored +/// as a vec of [`PeriodValue`]s that represents the aggregated value over a period of time for all +/// metrics. +#[derive(Default, Debug, Clone)] +pub struct PeriodicAggregatorState { + current_values: Option>>, +} + +impl PeriodicAggregatorState { + fn process_value( + &mut self, + value: PeriodValue, + agg_freq: &AggregationFrequency, + agg_func: &AggFuncF64, + ) -> Option> { + if let Some(current_values) = self.current_values.as_mut() { + // SAFETY: The current_values vector is guaranteed to contain at least one value. + let current_period_start = current_values + .first() + .expect("Aggregation state contains no values when at least one is expected.") + .start; + + // Determine if the value is in the current period + if agg_freq.is_date_in_period(¤t_period_start, &value.start) { + // New value in the current aggregation period; just append it. + current_values.push(value); + + None + } else { + // New value is part of a different period (assume the next one). + + // Calculate the aggregated value of the previous period. + let agg_period = if let Some(agg_value) = agg_func.calc_period_values(current_values) { + let agg_duration = value.start - current_period_start; + Some(PeriodValue::new(current_period_start, agg_duration.into(), agg_value)) + } else { + None + }; + + // Reset the state for the next period + current_values.clear(); + current_values.push(value); + + // Finally return the aggregated value from the previous period + agg_period + } + } else { + // No previous values defined; just append the value + self.current_values = Some(vec![value]); + + None + } + } + + fn process_value_no_period(&mut self, value: PeriodValue) { + if let Some(current_values) = self.current_values.as_mut() { + current_values.push(value); + } else { + self.current_values = Some(vec![value]); + } + } + + fn calc_aggregation(&self, agg_func: &AggFuncF64) -> Option> { + if let Some(current_values) = &self.current_values { + if let Some(agg_value) = agg_func.calc_period_values(current_values) { + // SAFETY: The current_values vector is guaranteed to contain at least one value. + let current_period_start = current_values + .first() + .expect("Aggregation state contains no values when at least one is expected.") + .start; + + let current_period_end = current_values + .last() + .expect("Aggregation state contains no values when at least one is expected.") + .start; + let current_period_duration = current_period_end - current_period_start; + Some(PeriodValue::new( + current_period_start, + current_period_duration.into(), + agg_value, + )) + } else { + None + } + } else { + None + } + } +} + +#[derive(Debug, Copy, Clone)] +pub struct PeriodValue { + pub start: NaiveDateTime, + pub duration: PywrDuration, + pub value: T, +} + +impl PeriodValue { + pub fn new(start: NaiveDateTime, duration: PywrDuration, value: T) -> Self { + Self { start, duration, value } + } + + /// The end of the period. + pub fn end(&self) -> NaiveDateTime { + self.duration + self.start + } +} + +impl PeriodValue> { + pub fn index(&self, index: usize) -> PeriodValue + where + T: Copy, + { + PeriodValue { + start: self.start, + duration: self.duration, + value: self.value[index], + } + } + pub fn len(&self) -> usize { + self.value.len() + } + + pub fn is_empty(&self) -> bool { + self.value.is_empty() + } +} + +impl From<&[PeriodValue]> for PeriodValue> +where + T: Copy, +{ + fn from(values: &[PeriodValue]) -> Self { + let start = values.first().expect("Empty vector of period values.").start; + let duration = values.last().expect("Empty vector of period values.").duration; + + let value = values.iter().map(|v| v.value).collect(); + Self { start, duration, value } + } +} + +#[derive(Clone, Debug)] +pub struct PeriodicAggregator { + frequency: Option, + function: AggFuncF64, +} + +impl PeriodicAggregator { + pub fn new(frequency: Option, function: AggFuncF64) -> Self { + Self { frequency, function } + } + + /// Append a new value to the aggregator. + /// + /// The new value should sequentially follow from the previously processed values. If the + /// value completes a new aggregation period then a value representing that aggregation is + /// returned. + pub fn process_value( + &self, + current_state: &mut PeriodicAggregatorState, + value: PeriodValue, + ) -> Option> { + // Split the given period into separate periods that align with the aggregation period. + let mut agg_value = None; + + if let Some(period) = &self.frequency { + for v in period.split_value_into_periods(value) { + let av = current_state.process_value(v, period, &self.function); + if av.is_some() { + if agg_value.is_some() { + panic!( + "Multiple aggregated values yielded from aggregator. This indicates that the given value spans multiple aggregation periods which is not supported." + ) + } + agg_value = av; + } + } + } else { + current_state.process_value_no_period(value); + } + agg_value + } + + pub fn calc_aggregation(&self, state: &PeriodicAggregatorState) -> Option> { + state.calc_aggregation(&self.function) + } + + /// Expected number of periods in the given time domain. + pub fn number_of_periods(&self, time_domain: &TimeDomain) -> usize { + match &self.frequency { + Some(frequency) => frequency.number_of_periods(time_domain), + None => 1, + } + } +} + +#[cfg(test)] +mod tests { + use super::{AggregationFrequency, PeriodicAggregator, PeriodicAggregatorState}; + use crate::agg_funcs::AggFuncF64; + use crate::recorders::aggregator::PeriodValue; + use chrono::{NaiveDate, TimeDelta}; + use float_cmp::assert_approx_eq; + + #[test] + fn test_periodic_aggregator() { + let agg = PeriodicAggregator { + frequency: Some(AggregationFrequency::Monthly), + function: AggFuncF64::Sum, + }; + + let mut state = PeriodicAggregatorState::default(); + + let start = NaiveDate::from_ymd_opt(2023, 1, 30) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); + assert!(agg_value.is_none()); + + let start = NaiveDate::from_ymd_opt(2023, 1, 31) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); + assert!(agg_value.is_none()); + + let start = NaiveDate::from_ymd_opt(2023, 2, 1) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); + assert!(agg_value.is_some()); + + let start = NaiveDate::from_ymd_opt(2023, 2, 2) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(); + let agg_value = agg.process_value(&mut state, PeriodValue::new(start, TimeDelta::days(1).into(), 1.0)); + assert!(agg_value.is_none()); + } + + #[test] + fn test_sub_daily_aggregation() { + let values = vec![ + PeriodValue::new( + NaiveDate::from_ymd_opt(2023, 1, 1) + .unwrap() + .and_hms_opt(0, 0, 0) + .unwrap(), + TimeDelta::hours(1).into(), + 2.0, + ), + PeriodValue::new( + NaiveDate::from_ymd_opt(2023, 1, 1) + .unwrap() + .and_hms_opt(1, 0, 0) + .unwrap(), + TimeDelta::hours(2).into(), + 1.0, + ), + PeriodValue::new( + NaiveDate::from_ymd_opt(2023, 1, 1) + .unwrap() + .and_hms_opt(3, 0, 0) + .unwrap(), + TimeDelta::hours(1).into(), + 3.0, + ), + ]; + + let agg_value = AggFuncF64::Mean.calc_period_values(values.as_slice()).unwrap(); + assert_approx_eq!(f64, agg_value, 7.0 / 4.0); + + let agg_value = AggFuncF64::Sum.calc_period_values(values.as_slice()).unwrap(); + let expected = 2.0 + 1.0 + 3.0; + assert_approx_eq!(f64, agg_value, expected); + } +} diff --git a/pywr-core/src/recorders/csv.rs b/pywr-core/src/recorders/csv.rs index 336d0d9e..f3f832cf 100644 --- a/pywr-core/src/recorders/csv.rs +++ b/pywr-core/src/recorders/csv.rs @@ -4,6 +4,7 @@ use super::{ }; use crate::models::ModelDomain; use crate::network::Network; +use crate::recorders::aggregator::AggregatorValue; use crate::recorders::metric_set::MetricSetIndex; use crate::scenario::ScenarioIndex; use crate::state::State; @@ -26,6 +27,8 @@ pub enum CsvError { #[source] source: ::csv::Error, }, + #[error("Event value is not supported in wide format")] + EventValueInWideFormat, } /// Output the values from a [`crate::recorders::MetricSet`] to a CSV file. @@ -61,15 +64,35 @@ impl CsvWideFmtOutput { index: self.metric_set_idx, })?; - if let Some(current_values) = metric_set_state.current_values() { - let values = current_values + // If the metric set has values then turn them into a row. + if metric_set_state.has_some_values() { + let values = metric_set_state + .current_values() .iter() - .map(|v| format!("{:.2}", v.value)) - .collect::>(); + .map(|maybe_v| match maybe_v { + Some(v) => match v { + AggregatorValue::Periodic(p) => Ok(format!("{:.2}", p.value)), + AggregatorValue::Event(_) => Err(CsvError::EventValueInWideFormat), + }, + None => Ok("".to_string()), // Missing value + }) + .collect::, _>>()?; // If the row is empty, add the start time if row.is_empty() { - row.push(current_values.first().unwrap().start.to_string()) + // Find the first non-None value and use that as the start time + let start = metric_set_state + .current_values() + .iter() + .find_map(|maybe_v| { + maybe_v.as_ref().and_then(|v| match v { + AggregatorValue::Periodic(p) => Some(p.start.to_string()), + AggregatorValue::Event(_) => None, + }) + }) + .unwrap_or_else(|| "unknown".to_string()); + + row.push(start) } row.extend(values); @@ -192,7 +215,7 @@ impl Recorder for CsvWideFmtOutput { } #[derive(Debug, Serialize, Deserialize)] -pub struct CsvLongFmtRecord { +pub struct CsvLongFmtValueRecord { time_start: NaiveDateTime, time_end: NaiveDateTime, simulation_id: usize, @@ -203,6 +226,17 @@ pub struct CsvLongFmtRecord { value: f64, } +#[derive(Debug, Serialize, Deserialize)] +pub struct CsvLongFmtEventRecord { + time_start: NaiveDateTime, + time_end: Option, + simulation_id: usize, + label: String, + metric_set: String, + name: String, + attribute: String, +} + /// Output the values from a several [`crate::recorders::MetricSet`]s to a CSV file in long format. /// /// The long format contains a row for each value produced by the metric set. This is useful @@ -245,37 +279,57 @@ impl CsvLongFmtOutput { .get(*metric_set_idx.deref()) .ok_or(CsvError::MetricSetIndexNotFound { index: *metric_set_idx })?; - if let Some(current_values) = metric_set_state.current_values() { - let metric_set = network - .get_metric_set(*metric_set_idx) - .ok_or(CsvError::MetricSetIndexNotFound { index: *metric_set_idx })?; + let metric_set = network + .get_metric_set(*metric_set_idx) + .ok_or(CsvError::MetricSetIndexNotFound { index: *metric_set_idx })?; - for (metric, value) in metric_set.iter_metrics().zip(current_values.iter()) { + for (metric, maybe_value) in metric_set.iter_metrics().zip(metric_set_state.current_values()) { + if let Some(value) = maybe_value { let name = metric.name().to_string(); let attribute = metric.attribute().to_string(); - let value_scaled = if let Some(decimal_places) = self.decimal_places { - let scale = 10.0_f64.powi(decimal_places.get() as i32); - (value.value * scale).round() / scale - } else { - value.value - }; - - let record = CsvLongFmtRecord { - time_start: value.start, - time_end: value.end(), - simulation_id: scenario_index.simulation_id(), - label: scenario_index.label(), - metric_set: metric_set.name().to_string(), - name, - attribute, - value: value_scaled, - }; - - internal.writer.serialize(record).map_err(|source| CsvError::CSVError { - path: self.filename.clone(), - source, - })?; + match value { + AggregatorValue::Periodic(value) => { + let value_scaled = if let Some(decimal_places) = self.decimal_places { + let scale = 10.0_f64.powi(decimal_places.get() as i32); + (value.value * scale).round() / scale + } else { + value.value + }; + + let record = CsvLongFmtValueRecord { + time_start: value.start, + time_end: value.end(), + simulation_id: scenario_index.simulation_id(), + label: scenario_index.label(), + metric_set: metric_set.name().to_string(), + name, + attribute, + value: value_scaled, + }; + + internal.writer.serialize(record).map_err(|source| CsvError::CSVError { + path: self.filename.clone(), + source, + })?; + } + AggregatorValue::Event(event) => { + let record = CsvLongFmtEventRecord { + time_start: event.start, + time_end: event.end, + simulation_id: scenario_index.simulation_id(), + label: scenario_index.label(), + metric_set: metric_set.name().to_string(), + name, + attribute, + }; + + internal.writer.serialize(record).map_err(|source| CsvError::CSVError { + path: self.filename.clone(), + source, + })?; + } + } } } } diff --git a/pywr-core/src/recorders/memory.rs b/pywr-core/src/recorders/memory.rs index ade8fcb5..9780a53a 100644 --- a/pywr-core/src/recorders/memory.rs +++ b/pywr-core/src/recorders/memory.rs @@ -1,7 +1,8 @@ use crate::agg_funcs::{AggFuncError, AggFuncF64}; use crate::models::ModelDomain; use crate::network::Network; -use crate::recorders::aggregator::PeriodValue; +use crate::recorders::aggregator::{AggregatorValue, Event, PeriodValue}; +use crate::recorders::metric_set::MetricSetOutputInfo; use crate::recorders::{ MetricSetIndex, MetricSetState, Recorder, RecorderAggregationError, RecorderDataFrameError, RecorderFinalResult, RecorderFinaliseError, RecorderInternalState, RecorderMeta, RecorderSaveError, RecorderSetupError, @@ -13,6 +14,7 @@ use crate::timestep::Timestep; use chrono::NaiveDateTime; use polars::df; use polars::frame::DataFrame; +use std::collections::HashMap; use std::ops::Deref; use thiserror::Error; use tracing::warn; @@ -25,6 +27,8 @@ pub enum AggregationError { AggregationFunctionFailed, #[error("Aggregation function error: {0}")] AggFuncError(#[from] AggFuncError), + #[error("Invalid aggregation order: {0}")] + InvalidOrder(String), } #[derive(Clone)] @@ -122,39 +126,39 @@ impl Aggregation { Ok(agg_value) } -} - -/// Internal state for the memory recorder. -/// -/// This is a 3D array, where the first dimension is the scenario, the second dimension is the time, -/// and the third dimension is the metric. -struct InternalState { - data: Vec>>>, -} - -impl InternalState { - fn new(num_scenarios: usize) -> Self { - let mut data: Vec>>> = Vec::with_capacity(num_scenarios); - for _ in 0..num_scenarios { - // We can't use `Vec::with_capacity` here because we don't know the number of - // periods that will be recorded. - data.push(Vec::new()) - } + /// Apply the time aggregation function to the provided events. + fn apply_time_func_events(&self, events: &[Event]) -> Result { + let agg_value = if events.len() == 1 { + if self.time.is_some() { + warn!("Aggregation function defined for time, but not used.") + } + events + .first() + .expect("No events found in time series") + .duration() + .map(|d| d.fractional_days()) + .ok_or(AggregationError::AggregationFunctionFailed)? + } else { + self.time + .as_ref() + .ok_or(AggregationError::AggregationFunctionNotDefined)? + .calc_events(events) + .ok_or(AggregationError::AggregationFunctionFailed)? + }; - Self { data } + Ok(agg_value) } } -struct LongFmtRecord { - time_start: NaiveDateTime, - time_end: NaiveDateTime, - simulation_id: usize, - label: String, - metric_set: String, - name: String, - attribute: String, - value: f64, +/// Periodic internal state for the memory recorder. +/// +/// This is a 3D array, where the first dimension is the scenario, the second dimension is the time, +/// and the third dimension is the metric. It is used for storing periodic output data which +/// produces a value for every scenario at the same time. +#[derive(Clone)] +struct PeriodicInternalState { + data: Vec>>>, } /// Final results for the memory recorder. @@ -166,7 +170,7 @@ pub struct MemoryRecorderResult { scenario_indices: Vec, metric_names: Vec, metric_attrs: Vec, - data: Vec>>>, + data: InternalState, aggregation: Aggregation, order: AggregationOrder, } @@ -176,72 +180,60 @@ impl MemoryRecorderResult { /// /// This method will first aggregation over the metrics, then over time, and finally over the scenarios. fn aggregate_metric_time_scenario(&self) -> Result { - let scenario_data: Vec = self - .data - .iter() - .map(|time_data| { - // Aggregate each metric at each time step; - // this results in a time series iterator of aggregated values - let ts: Vec> = time_data + match &self.data { + InternalState::Events(_) => Err(AggregationError::InvalidOrder( + "Cannot aggregate over events by metric first. Events must be aggregated by time first.".to_string(), + )), + InternalState::Periodic(state) => { + let scenario_data: Vec = state + .data .iter() - .map(|metric_data| self.aggregation.apply_metric_func_period_value(metric_data)) + .map(|time_data| { + // Aggregate each metric at each time step; + // this results in a time series iterator of aggregated values + let ts: Vec> = time_data + .iter() + .map(|metric_data| self.aggregation.apply_metric_func_period_value(metric_data)) + .collect::>()?; + + self.aggregation.apply_time_func(&ts) + }) .collect::>()?; - self.aggregation.apply_time_func(&ts) - }) - .collect::>()?; - - self.aggregation.apply_scenario_func(&scenario_data) + self.aggregation.apply_scenario_func(&scenario_data) + } + } } /// Aggregate over the saved data to a single value using the provided aggregation functions. /// /// This method will first aggregation over time, then over the metrics, and finally over the scenarios. fn aggregate_time_metric_scenario(&self) -> Result { - let scenario_data: Vec = self - .data - .iter() - .map(|time_data| { - // We expect the same number of metrics in all the entries - let num_metrics = time_data.first().expect("No metrics found in time data").len(); - - // Aggregate each metric over time first. This requires transposing the saved data. - let metric_ts: Vec = (0..num_metrics) - // TODO remove the collect allocation; requires `AggregationFunction.calc` to accept an iterator - .map(|metric_idx| time_data.iter().map(|t| t.index(metric_idx)).collect()) - .map(|ts: Vec>| self.aggregation.apply_time_func(&ts)) + match &self.data { + InternalState::Events(state) => state.aggregate_time_metric_scenario(&self.aggregation), + InternalState::Periodic(state) => { + let scenario_data: Vec = state + .data + .iter() + .map(|time_data| { + // We expect the same number of metrics in all the entries + let num_metrics = time_data.first().expect("No metrics found in time data").len(); + + // Aggregate each metric over time first. This requires transposing the saved data. + let metric_ts: Vec = (0..num_metrics) + // TODO remove the collect allocation; requires `AggregationFunction.calc` to accept an iterator + .map(|metric_idx| time_data.iter().map(|t| t.index(metric_idx)).collect()) + .map(|ts: Vec>| self.aggregation.apply_time_func(&ts)) + .collect::>()?; + + // Now aggregate over the metrics + self.aggregation.apply_metric_func_f64(&metric_ts) + }) .collect::>()?; - // Now aggregate over the metrics - self.aggregation.apply_metric_func_f64(&metric_ts) - }) - .collect::>()?; - - self.aggregation.apply_scenario_func(&scenario_data) - } - - fn iter_long_fmt_records(&self) -> impl Iterator { - self.scenario_indices - .iter() - .zip(self.data.iter()) - .flat_map(|(scenario_index, scenario_data)| { - scenario_data.iter().flat_map(|pv| { - pv.value - .iter() - .zip(self.metric_names.iter()) - .zip(self.metric_attrs.iter()) - .map(|((v, name), attr)| LongFmtRecord { - time_start: pv.start, - time_end: pv.end(), - simulation_id: scenario_index.simulation_id(), - label: scenario_index.label(), - metric_set: self.meta.name.clone(), - name: name.clone(), - attribute: attr.clone(), - value: *v, - }) - }) - }) + self.aggregation.apply_scenario_func(&scenario_data) + } + } } } @@ -262,22 +254,247 @@ impl RecorderFinalResult for MemoryRecorderResult { } fn to_dataframe(&self) -> Result { - let records: Vec = self.iter_long_fmt_records().collect(); - - df!( - "time_start" => records.iter().map(|r| r.time_start).collect::>(), - "time_end" => records.iter().map(|r| r.time_end).collect::>(), - "simulation_id" => records.iter().map(|r| r.simulation_id as u32).collect::>(), - "label" => records.iter().map(|r| r.label.as_str()).collect::>(), - "metric_set" => records.iter().map(|r| r.metric_set.as_str()).collect::>(), - "name" => records.iter().map(|r| r.name.as_str()).collect::>(), - "attribute" => records.iter().map(|r| r.attribute.as_str()).collect::>(), - "value" => records.iter().map(|r| r.value).collect::>(), - ) - .map_err(|source| RecorderDataFrameError::PolarsError { - name: self.meta.name.clone(), - source, - }) + match &self.data { + InternalState::Events(state) => { + let mut time_start = Vec::new(); + let mut time_end = Vec::new(); + let mut simulation_id = Vec::new(); + let mut label = Vec::new(); + let mut metric_set = Vec::new(); + let mut names = Vec::new(); + let mut attribute = Vec::new(); + + self.scenario_indices + .iter() + .zip(state.events.iter()) + .for_each(|(scenario_index, scenario_events)| { + scenario_events.iter().for_each(|ev| { + let name = &self.metric_names[ev.metric_index]; + let attr = &self.metric_attrs[ev.metric_index]; + + time_start.push(ev.start); + time_end.push(ev.end); + simulation_id.push(scenario_index.simulation_id() as u32); + label.push(scenario_index.label()); + metric_set.push(self.meta.name.clone()); + names.push(name.clone()); + attribute.push(attr.clone()); + }) + }); + + df!( + "time_start" => time_start, + "time_end" => time_end, + "simulation_id" => simulation_id, + "label" => label, + "metric_set" => metric_set, + "name" => names, + "attribute" => attribute, + ) + .map_err(|source| RecorderDataFrameError::PolarsError { + name: self.meta.name.clone(), + source, + }) + } + InternalState::Periodic(state) => { + let mut time_start = Vec::new(); + let mut time_end = Vec::new(); + let mut simulation_id = Vec::new(); + let mut label = Vec::new(); + let mut metric_set = Vec::new(); + let mut names = Vec::new(); + let mut attribute = Vec::new(); + let mut value = Vec::new(); + + self.scenario_indices + .iter() + .zip(state.data.iter()) + .for_each(|(scenario_index, scenario_data)| { + scenario_data.iter().for_each(|pv| { + pv.value + .iter() + .zip(self.metric_names.iter()) + .zip(self.metric_attrs.iter()) + .for_each(|((v, name), attr)| { + time_start.push(pv.start); + time_end.push(pv.end()); + simulation_id.push(scenario_index.simulation_id() as u32); + label.push(scenario_index.label()); + metric_set.push(self.meta.name.clone()); + names.push(name.clone()); + attribute.push(attr.clone()); + value.push(*v); + }) + }) + }); + + df!( + "time_start" => time_start, + "time_end" => time_end, + "simulation_id" => simulation_id, + "label" => label, + "metric_set" => metric_set, + "name" => names, + "attribute" => attribute, + "value" => value, + ) + .map_err(|source| RecorderDataFrameError::PolarsError { + name: self.meta.name.clone(), + source, + }) + } + } + } +} + +#[derive(Copy, Clone)] +struct MemoryEvent { + start: NaiveDateTime, + end: Option, + metric_index: usize, +} + +impl MemoryEvent { + fn from_event(event: Event, metric_index: usize) -> MemoryEvent { + MemoryEvent { + start: event.start, + end: event.end, + metric_index, + } + } +} + +impl From for Event { + fn from(me: MemoryEvent) -> Self { + Event { + start: me.start, + end: me.end, + } + } +} + +/// Event internal state for the memory recorder. +/// +/// This is a nested vector of events where the outer vec is the length of the scenarios, +/// and the inner vector are the events for that scenario. +#[derive(Clone)] +struct EventInternalState { + events: Vec>, +} + +impl EventInternalState { + /// Aggregate over the saved data to a single value using the provided aggregation functions. + /// + /// This method will first aggregation over time, then over the metrics, and finally over the scenarios. + fn aggregate_time_metric_scenario(&self, aggregation: &Aggregation) -> Result { + let scenario_data: Vec = self + .events + .iter() + .map(|events| { + // Accumulate the events for each metric + let mut events_by_metric: HashMap> = HashMap::new(); + + for event in events { + events_by_metric + .entry(event.metric_index) + .or_default() + .push((*event).into()); + } + + // Aggregate each metric over time first. + // NB, these are not necessarily in order of the metric index. + // Some metrics may not have any events. + let metric_ts: Vec = events_by_metric + .values() + .map(|metric_events| aggregation.apply_time_func_events(metric_events)) + .collect::>()?; + + // Now aggregate over the metrics + aggregation.apply_metric_func_f64(&metric_ts) + }) + .collect::>()?; + + aggregation.apply_scenario_func(&scenario_data) + } +} + +/// Internal state for the memory recorder. +/// +/// The variant used depends on the type of data produced by the aggregator. +#[derive(Clone)] +enum InternalState { + Periodic(PeriodicInternalState), + Events(EventInternalState), +} + +impl InternalState { + fn new_periodic(num_scenarios: usize, num_periods: Option) -> Self { + let mut data: Vec>>> = Vec::with_capacity(num_scenarios); + + for _ in 0..num_scenarios { + data.push(Vec::with_capacity(num_periods.unwrap_or_default())) + } + + Self::Periodic(PeriodicInternalState { data }) + } + + fn new_event(num_scenarios: usize) -> Self { + let mut events: Vec<_> = Vec::with_capacity(num_scenarios); + + for _ in 0..num_scenarios { + events.push(Vec::new()); + } + + Self::Events(EventInternalState { events }) + } + + fn append_value(&mut self, scenario_index: &ScenarioIndex, values: &[Option]) { + match self { + Self::Periodic(state) => { + let scenario_data = state + .data + .get_mut(scenario_index.simulation_id()) + .expect("No scenario data found"); + + // Find the first non-None value and use that as the start time + let (start, duration) = values + .iter() + .find_map(|maybe_v| { + maybe_v.as_ref().and_then(|v| match v { + AggregatorValue::Periodic(p) => Some((p.start, p.duration)), + AggregatorValue::Event(_) => None, + }) + }) + .unwrap_or_else(|| panic!("Could not determine time-step information.")); + + let period_values = values + .iter() + .map(|maybe_v| match maybe_v { + Some(v) => match v { + AggregatorValue::Periodic(v) => v.value, + AggregatorValue::Event(_) => panic!("Cannot append event values to periodic data."), + }, + None => panic!("No value found for metric."), + }) + .collect::>(); + + scenario_data.push(PeriodValue::new(start, duration, period_values)); + } + Self::Events(state) => { + let scenario_data = state + .events + .get_mut(scenario_index.simulation_id()) + .expect("No scenario data found"); + + for (metric_idx, value) in values.iter().enumerate() { + match value { + Some(AggregatorValue::Event(e)) => scenario_data.push(MemoryEvent::from_event(*e, metric_idx)), + Some(AggregatorValue::Periodic(_)) => panic!("Cannot append periodic values to event data."), + None => panic!("No value found for metric."), + } + } + } + } } } @@ -322,17 +539,29 @@ impl Recorder for MemoryRecorder { fn setup( &self, domain: &ModelDomain, - _network: &Network, + network: &Network, ) -> Result>, RecorderSetupError> { - let data = InternalState::new(domain.scenarios().len()); + let metric_set = + network + .get_metric_set(self.metric_set_idx) + .ok_or_else(|| RecorderSetupError::MetricSetIndexNotFound { + index: self.metric_set_idx, + })?; + + let state = match metric_set.output_info(domain.time()) { + MetricSetOutputInfo::Periodic { num_periods } => { + InternalState::new_periodic(domain.scenarios().len(), Some(num_periods)) + } + MetricSetOutputInfo::Event => InternalState::new_event(domain.scenarios().len()), + }; - Ok(Some(Box::new(data))) + Ok(Some(Box::new(state))) } fn save( &self, _timestep: &Timestep, - _scenario_indices: &[ScenarioIndex], + scenario_indices: &[ScenarioIndex], _model: &Network, _state: &[State], metric_set_states: &[Vec], @@ -340,16 +569,16 @@ impl Recorder for MemoryRecorder { ) -> Result<(), RecorderSaveError> { let internal_state = downcast_internal_state_mut::(internal_state); - // Iterate through all of the scenario's state - for (ms_scenario_states, scenario_data) in metric_set_states.iter().zip(internal_state.data.iter_mut()) { + // Iterate through all the scenario's state + for (scenario_index, ms_scenario_states) in scenario_indices.iter().zip(metric_set_states.iter()) { let metric_set_state = ms_scenario_states.get(*self.metric_set_idx.deref()).ok_or_else(|| { RecorderSaveError::MetricSetIndexNotFound { index: self.metric_set_idx, } })?; - if let Some(current_values) = metric_set_state.current_values() { - scenario_data.push(current_values.into()); + if metric_set_state.has_some_values() { + internal_state.append_value(scenario_index, metric_set_state.current_values()); } } @@ -372,16 +601,16 @@ impl Recorder for MemoryRecorder { index: self.metric_set_idx, })?; - // Iterate through all of the scenario's state - for (ms_scenario_states, scenario_data) in metric_set_states.iter().zip(internal_state.data.iter_mut()) { + // Iterate through all the scenario's state + for (scenario_index, ms_scenario_states) in scenario_indices.iter().zip(metric_set_states.iter()) { let metric_set_state = ms_scenario_states.get(*self.metric_set_idx.deref()).ok_or_else(|| { RecorderFinaliseError::MetricSetIndexNotFound { index: self.metric_set_idx, } })?; - if let Some(current_values) = metric_set_state.current_values() { - scenario_data.push(current_values.into()); + if metric_set_state.has_some_values() { + internal_state.append_value(scenario_index, metric_set_state.current_values()); } } @@ -390,7 +619,7 @@ impl Recorder for MemoryRecorder { scenario_indices: scenario_indices.to_vec(), metric_names: metric_set.iter_metrics().map(|m| m.name().to_string()).collect(), metric_attrs: metric_set.iter_metrics().map(|m| m.attribute().to_string()).collect(), - data: internal_state.data, + data: internal_state.deref().clone(), aggregation: self.aggregation.clone(), order: self.order, }; @@ -405,9 +634,10 @@ mod tests { use crate::agg_funcs::AggFuncF64; use crate::models::ModelDomain; use crate::recorders::RecorderMeta; - use crate::recorders::aggregator::PeriodValue; + use crate::recorders::aggregator::{AggregatorValue, Event, PeriodValue}; use crate::scenario::{ScenarioDomainBuilder, ScenarioGroupBuilder}; use crate::test_utils::default_timestepper; + use chrono::NaiveDate; use float_cmp::assert_approx_eq; use rand::{Rng, SeedableRng}; use rand_chacha::ChaCha8Rng; @@ -422,7 +652,7 @@ mod tests { let domain = ModelDomain::try_from(default_timestepper(), scenario_builder).unwrap(); let num_metrics = 3; - let mut state = InternalState::new(domain.scenarios().len()); + let mut state = InternalState::new_periodic(domain.scenarios().len(), None); let mut rng = ChaCha8Rng::seed_from_u64(0); let dist: Normal = Normal::new(0.0, 1.0).unwrap(); @@ -432,24 +662,26 @@ mod tests { let mut count_non_zero_by_metric = vec![0.0; num_metrics]; domain.time().timesteps().iter().for_each(|timestep| { - state.data.iter_mut().for_each(|scenario_data| { - let metric_data = (&mut rng).sample_iter(&dist).take(num_metrics).collect::>(); + if let InternalState::Periodic(state) = &mut state { + state.data.iter_mut().for_each(|scenario_data| { + let metric_data = (&mut rng).sample_iter(&dist).take(num_metrics).collect::>(); - // Compute the expected values - if metric_data.iter().sum::() > 0.0 { - count_non_zero_max += 1.0; - } - // ... and by metric - metric_data.iter().enumerate().for_each(|(i, v)| { - if *v > 0.0 { - count_non_zero_by_metric[i] += 1.0; + // Compute the expected values + if metric_data.iter().sum::() > 0.0 { + count_non_zero_max += 1.0; } - }); + // ... and by metric + metric_data.iter().enumerate().for_each(|(i, v)| { + if *v > 0.0 { + count_non_zero_by_metric[i] += 1.0; + } + }); - let metric_data = PeriodValue::new(timestep.date, timestep.duration, metric_data); + let metric_data = PeriodValue::new(timestep.date, timestep.duration, metric_data); - scenario_data.push(metric_data); - }); + scenario_data.push(metric_data); + }); + } }); let agg = Aggregation::new( @@ -463,7 +695,7 @@ mod tests { scenario_indices: domain.scenarios().indices().to_vec(), metric_names: vec!["m1".to_string(), "m2".to_string(), "m3".to_string()], metric_attrs: vec!["a1".to_string(), "a2".to_string(), "a3".to_string()], - data: state.data, + data: state.clone(), aggregation: agg, order: super::AggregationOrder::MetricTimeScenario, }; @@ -474,4 +706,53 @@ mod tests { let agg_value = result.aggregate_time_metric_scenario().expect("Aggregation failed"); assert_approx_eq!(f64, agg_value, count_non_zero_by_metric.iter().sum()); } + + #[test] + fn test_memory_event_aggregation() { + let mut scenario_builder = ScenarioDomainBuilder::default(); + let scenario_group = ScenarioGroupBuilder::new("test-scenario", 2).build().unwrap(); + scenario_builder = scenario_builder.with_group(scenario_group).unwrap(); + + let domain = ModelDomain::try_from(default_timestepper(), scenario_builder).unwrap(); + + let num_metrics = 3; + let mut state = InternalState::new_event(domain.scenarios().len()); + + for scenario_index in domain.scenarios().indices() { + for event_index in 0..4 { + // Create an event with a known start and end time + let start = NaiveDate::from_ymd_opt(2016, event_index + 1, 8).unwrap(); + let end = start + chrono::Duration::days(event_index as i64 + 1); + + let events: Vec<_> = (0..num_metrics) + .map(|_| { + let e = Event { + start: start.into(), + end: Some(end.into()), + }; + Some(AggregatorValue::Event(e)) + }) + .collect(); + + state.append_value(scenario_index, &events); + } + } + + // This should be the total duration of all the events + let agg = Aggregation::new(Some(AggFuncF64::Sum), Some(AggFuncF64::Sum), Some(AggFuncF64::Sum)); + + let result = MemoryRecorderResult { + meta: RecorderMeta::new("test"), + scenario_indices: domain.scenarios().indices().to_vec(), + metric_names: vec!["m1".to_string(), "m2".to_string(), "m3".to_string()], + metric_attrs: vec!["a1".to_string(), "a2".to_string(), "a3".to_string()], + data: state.clone(), + aggregation: agg, + order: super::AggregationOrder::MetricTimeScenario, + }; + + let expected_total_duration = domain.scenarios().len() as f64 * num_metrics as f64 * (1.0 + 2.0 + 3.0 + 4.0); + let agg_value = result.aggregate_time_metric_scenario().expect("Aggregation failed"); + assert_approx_eq!(f64, agg_value, expected_total_duration); + } } diff --git a/pywr-core/src/recorders/metric_set.rs b/pywr-core/src/recorders/metric_set.rs index 54630108..f79a1a93 100644 --- a/pywr-core/src/recorders/metric_set.rs +++ b/pywr-core/src/recorders/metric_set.rs @@ -1,9 +1,9 @@ use crate::metric::{MetricF64, MetricF64Error}; use crate::network::Network; -use crate::recorders::aggregator::{Aggregator, AggregatorState, PeriodValue}; +use crate::recorders::aggregator::{AggregatorValue, NestedAggregator, NestedAggregatorState, PeriodValue}; use crate::scenario::ScenarioIndex; use crate::state::State; -use crate::timestep::Timestep; +use crate::timestep::{TimeDomain, Timestep}; use std::fmt; use std::fmt::{Display, Formatter}; use std::ops::Deref; @@ -80,18 +80,35 @@ impl Display for MetricSetIndex { #[derive(Debug, Clone)] pub struct MetricSetState { - // Populated with any yielded values from the last processing. - current_values: Option>>, - // If the metric set aggregates then this state tracks the aggregation of each metric - aggregation_states: Option>, + /// Populated with any yielded values from the last processing. One entry per + /// metric in the set. + current_values: Vec>, + /// If the metric set aggregates then this state tracks the aggregation of each metric + aggregation_states: Option>, } impl MetricSetState { - pub fn current_values(&self) -> Option<&[PeriodValue]> { - self.current_values.as_deref() + /// Returns the current values for the metrics in the set. There is an entry for each metric + /// in the set, which will be `None` if no value was yielded for that metric. + pub fn current_values(&self) -> &[Option] { + self.current_values.as_slice() + } + + /// Helper method to determine if there are any values in the current state. + pub fn has_some_values(&self) -> bool { + self.current_values.iter().any(|v| v.is_some()) } } +/// Information about the type of output expected from a [`MetricSet`]. +pub enum MetricSetOutputInfo { + Periodic { + // The number of time periods expected in the output + num_periods: usize, + }, + Event, +} + #[derive(Debug, Error)] pub enum MetricSetSaveError { #[error("Metric error: {0}")] @@ -102,12 +119,12 @@ pub enum MetricSetSaveError { #[derive(Clone, Debug)] pub struct MetricSet { name: String, - aggregator: Option, + aggregator: Option, metrics: Vec, } impl MetricSet { - pub fn new(name: &str, aggregator: Option, metrics: Vec) -> Self { + pub fn new(name: &str, aggregator: Option, metrics: Vec) -> Self { Self { name: name.to_string(), aggregator, @@ -126,7 +143,7 @@ impl MetricSet { /// Setup a new [`MetricSetState`] for this [`MetricSet`]. pub fn setup(&self) -> MetricSetState { MetricSetState { - current_values: None, + current_values: vec![None; self.metrics.len()], aggregation_states: self .aggregator .as_ref() @@ -134,6 +151,16 @@ impl MetricSet { } } + pub fn output_info(&self, time_domain: &TimeDomain) -> MetricSetOutputInfo { + match &self.aggregator { + Some(aggregator) => aggregator.output_info(time_domain), + None => MetricSetOutputInfo::Periodic { + // Without an aggregator the output will be on per time-step. + num_periods: time_domain.len(), + }, + } + } + pub fn save( &self, timestep: &Timestep, @@ -168,24 +195,14 @@ impl MetricSet { // Use a for loop instead of using an iterator because we need to execute the // `append_value` method on all aggregators. for (value, current_state) in values.iter().zip(aggregation_states.iter_mut()) { - if let Some(agg_value) = aggregator.append_value(current_state, *value) { - agg_values.push(agg_value); - } - } + let agg_value = (*value).into(); - let agg_values = if agg_values.is_empty() { - None - } else if agg_values.len() == values.len() { - Some(agg_values) - } else { - // This should never happen because the aggregator should either yield no values - // or the same number of values as the input metrics. - unreachable!("Some values were aggregated and some were not!"); - }; + agg_values.push(aggregator.append_value(current_state, agg_value)); + } internal_state.current_values = agg_values; } else { - internal_state.current_values = Some(values); + internal_state.current_values = values.into_iter().map(|v| Some(v.into())).collect(); } Ok(()) @@ -201,11 +218,11 @@ impl MetricSet { let final_values = aggregation_states .iter_mut() .map(|current_state| aggregator.finalise(current_state)) - .collect::>>(); + .collect::>(); internal_state.current_values = final_values; } else { - internal_state.current_values = None; + internal_state.current_values = vec![None; self.metrics.len()]; } } } diff --git a/pywr-core/src/recorders/mod.rs b/pywr-core/src/recorders/mod.rs index 15c6e7aa..a266f243 100644 --- a/pywr-core/src/recorders/mod.rs +++ b/pywr-core/src/recorders/mod.rs @@ -16,8 +16,10 @@ use crate::recorders::hdf::Hdf5Error; use crate::scenario::ScenarioIndex; use crate::state::State; use crate::timestep::Timestep; -pub use aggregator::{AggregationFrequency, Aggregator, PeriodValue}; -pub use csv::{CsvLongFmtOutput, CsvLongFmtRecord, CsvWideFmtOutput}; +pub use aggregator::{ + AggregationFrequency, Aggregator, Event, EventAggregator, NestedAggregator, PeriodValue, PeriodicAggregator, +}; +pub use csv::{CsvLongFmtOutput, CsvWideFmtOutput}; use float_cmp::{ApproxEq, F64Margin, approx_eq}; #[cfg(feature = "hdf5")] pub use hdf::HDF5Recorder; diff --git a/pywr-core/src/test_utils.rs b/pywr-core/src/test_utils.rs index b6185240..539b93a9 100644 --- a/pywr-core/src/test_utils.rs +++ b/pywr-core/src/test_utils.rs @@ -7,7 +7,7 @@ use crate::network::{Network, NetworkError}; use crate::node::StorageInitialVolume; use crate::parameters::{AggregatedParameter, Array2Parameter, ConstantParameter, GeneralParameter}; use crate::recorders::{AssertionF64Recorder, AssertionU64Recorder}; -use crate::scenario::{ScenarioDomainBuilder, ScenarioGroupBuilder}; +use crate::scenario::{ScenarioDomain, ScenarioDomainBuilder, ScenarioGroupBuilder}; #[cfg(feature = "cbc")] use crate::solvers::CbcSolver; #[cfg(feature = "ipm-ocl")] @@ -68,6 +68,17 @@ pub fn default_model() -> Model { Model::new(domain, network) } +/// Create a test scenario domain with a single scenario group containing the specified number of scenarios. +pub fn test_scenario_domain(num_scenarios: usize) -> ScenarioDomain { + let mut scenario_builder = ScenarioDomainBuilder::default(); + let scenario_group = ScenarioGroupBuilder::new("test-scenario", num_scenarios) + .build() + .unwrap(); + scenario_builder = scenario_builder.with_group(scenario_group).unwrap(); + + scenario_builder.build().expect("Failed to build Scenario domain.") +} + /// Create a simple test network with three nodes. pub fn simple_network(network: &mut Network, inflow_scenario_index: usize, num_inflow_scenarios: usize) { let input_node = network.add_input_node("input", None).unwrap(); diff --git a/pywr-core/src/timestep.rs b/pywr-core/src/timestep.rs index d384ff20..b86315a6 100644 --- a/pywr-core/src/timestep.rs +++ b/pywr-core/src/timestep.rs @@ -393,11 +393,11 @@ impl TimeDomain { self.timesteps.len() } - pub fn first_timestep(&self) -> Option<&Timestep> { + pub fn first(&self) -> Option<&Timestep> { self.timesteps.first() } - pub fn last_timestep(&self) -> Option<&Timestep> { + pub fn last(&self) -> Option<&Timestep> { self.timesteps.last() } diff --git a/pywr-schema/src/lib.rs b/pywr-schema/src/lib.rs index 123de730..78228010 100644 --- a/pywr-schema/src/lib.rs +++ b/pywr-schema/src/lib.rs @@ -17,6 +17,7 @@ pub mod nodes; pub mod outputs; pub mod parameters; mod py_utils; +mod predicate; pub mod timeseries; mod v1; mod visit; diff --git a/pywr-schema/src/metric_sets/mod.rs b/pywr-schema/src/metric_sets/mod.rs index f0561229..ed77d314 100644 --- a/pywr-schema/src/metric_sets/mod.rs +++ b/pywr-schema/src/metric_sets/mod.rs @@ -8,6 +8,7 @@ use crate::metric::{EdgeReference, VirtualNodeAttrReference}; use crate::network::LoadArgs; #[cfg(feature = "core")] use crate::parameters::{Parameter, PythonReturnType}; +use crate::predicate::Predicate; use pywr_schema_macros::skip_serializing_none; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; @@ -37,18 +38,15 @@ impl From for pywr_core::recorders::AggregationFrequency { } } -/// A set of metrics that can be output from a model run. +/// Periodic aggregation of metric values. /// -/// A metric set can optionally have an aggregator, which will apply an aggregation function -/// over the metrics in the set. If an aggregation frequency is provided then the aggregation -/// will be performed over each period implied by that frequency. For example, if the frequency -/// is monthly then the aggregation will be performed over each month in the model run. +/// Applies an aggregation function over metric values at a specified frequency. If +/// no frequency is specified, the aggregation is applied over all values. /// -/// If the metric set has a child aggregator then the aggregation will be performed over the -/// aggregated values of the child aggregator. +/// An optional child aggregator can be specified to allow for nested aggregations. #[derive(Deserialize, Serialize, Clone, JsonSchema)] #[serde(deny_unknown_fields)] -pub struct MetricAggregator { +pub struct PeriodicMetricAggregator { /// Optional aggregation frequency. pub freq: Option, /// Aggregation function to apply over metric values. @@ -58,15 +56,60 @@ pub struct MetricAggregator { } #[cfg(feature = "core")] -impl MetricAggregator { +impl PeriodicMetricAggregator { fn load(&self, data_path: Option<&Path>) -> Result { - let child = self.child.as_ref().map(|a| a.load(data_path)).transpose()?; + Ok( + pywr_core::recorders::PeriodicAggregator::new(self.freq.map(|p| p.into()), self.func.load(data_path)?) + .into(), + ) + } +} + +/// Event-based aggregation of metric values. +/// +/// Starts a new event when the `predicate` is true relative to the `threshold`. The event +/// continues until the `predicate` is false. +/// +/// An optional child aggregator can be specified to allow for nested aggregations. +#[derive(Deserialize, Serialize, Clone, JsonSchema)] +#[serde(deny_unknown_fields)] +pub struct EventMetricAggregator { + pub predicate: Predicate, + pub threshold: f64, + /// Optional child aggregator. + pub child: Option>, +} + +#[cfg(feature = "core")] +impl EventMetricAggregator { + fn load(&self, _data_path: Option<&Path>) -> Result { + let ema = pywr_core::recorders::EventAggregator::new(self.predicate.into(), self.threshold); + Ok(ema.into()) + } +} + +#[derive(Deserialize, Serialize, Clone, JsonSchema)] +#[serde(tag = "type")] +pub enum MetricAggregator { + Periodic(PeriodicMetricAggregator), + Event(EventMetricAggregator), +} + +#[cfg(feature = "core")] +impl MetricAggregator { + fn load(&self, data_path: Option<&Path>) -> Result { + let (agg, child) = match self { + MetricAggregator::Periodic(p) => ( + p.load(data_path)?, + p.child.as_ref().map(|c| c.load(data_path)).transpose()?, + ), + MetricAggregator::Event(e) => ( + e.load(data_path)?, + e.child.as_ref().map(|c| c.load(data_path)).transpose()?, + ), + }; - Ok(pywr_core::recorders::Aggregator::new( - self.freq.map(|p| p.into()), - self.func.load(data_path)?, - child, - )) + Ok(pywr_core::recorders::NestedAggregator::new(agg, child)) } } diff --git a/pywr-schema/src/parameters/mod.rs b/pywr-schema/src/parameters/mod.rs index a7995613..8c3b47d2 100644 --- a/pywr-schema/src/parameters/mod.rs +++ b/pywr-schema/src/parameters/mod.rs @@ -71,7 +71,7 @@ use schemars::JsonSchema; use std::path::{Path, PathBuf}; use strum_macros::{Display, EnumDiscriminants, EnumIter, EnumString, IntoStaticStr}; pub use tables::TablesArrayParameter; -pub use thresholds::{MultiThresholdParameter, Predicate, ThresholdParameter}; +pub use thresholds::{MultiThresholdParameter, ThresholdParameter}; #[skip_serializing_none] #[derive(serde::Deserialize, serde::Serialize, Debug, Clone, JsonSchema, PywrVisitAll)] diff --git a/pywr-schema/src/parameters/thresholds.rs b/pywr-schema/src/parameters/thresholds.rs index e2b6c9eb..cf3054c7 100644 --- a/pywr-schema/src/parameters/thresholds.rs +++ b/pywr-schema/src/parameters/thresholds.rs @@ -6,6 +6,7 @@ use crate::metric::{Metric, NodeAttrReference}; #[cfg(feature = "core")] use crate::network::LoadArgs; use crate::parameters::{ConversionData, ParameterMeta}; +use crate::predicate::Predicate; use crate::v1::{IntoV2, TryFromV1, try_convert_parameter_attr}; #[cfg(feature = "core")] use pywr_core::parameters::{ParameterName, ParameterType}; @@ -14,49 +15,9 @@ use pywr_v1_schema::parameters::{ MultipleThresholdIndexParameter as MultiThresholdIndexParameterV1, MultipleThresholdParameterIndexParameter as MultipleThresholdParameterIndexParameterV1, NodeThresholdParameter as NodeThresholdParameterV1, ParameterThresholdParameter as ParameterThresholdParameterV1, - Predicate as PredicateV1, StorageThresholdParameter as StorageThresholdParameterV1, + StorageThresholdParameter as StorageThresholdParameterV1, }; use schemars::JsonSchema; -use strum_macros::{Display, EnumIter}; - -#[derive(serde::Deserialize, serde::Serialize, Debug, Clone, Copy, JsonSchema, PywrVisitAll, Display, EnumIter)] -pub enum Predicate { - #[serde(alias = "<")] - LT, - #[serde(alias = ">")] - GT, - #[serde(alias = "==")] - EQ, - #[serde(alias = "<=")] - LE, - #[serde(alias = ">=")] - GE, -} - -impl From for Predicate { - fn from(v1: PredicateV1) -> Self { - match v1 { - PredicateV1::LT => Predicate::LT, - PredicateV1::GT => Predicate::GT, - PredicateV1::EQ => Predicate::EQ, - PredicateV1::LE => Predicate::LE, - PredicateV1::GE => Predicate::GE, - } - } -} - -#[cfg(feature = "core")] -impl From for pywr_core::parameters::Predicate { - fn from(p: Predicate) -> Self { - match p { - Predicate::LT => pywr_core::parameters::Predicate::LessThan, - Predicate::GT => pywr_core::parameters::Predicate::GreaterThan, - Predicate::EQ => pywr_core::parameters::Predicate::EqualTo, - Predicate::LE => pywr_core::parameters::Predicate::LessThanOrEqualTo, - Predicate::GE => pywr_core::parameters::Predicate::GreaterThanOrEqualTo, - } - } -} /// A parameter that compares a metric against a threshold metric /// diff --git a/pywr-schema/src/predicate.rs b/pywr-schema/src/predicate.rs new file mode 100644 index 00000000..c07d79e9 --- /dev/null +++ b/pywr-schema/src/predicate.rs @@ -0,0 +1,43 @@ +use pywr_schema_macros::PywrVisitAll; +use pywr_v1_schema::parameters::Predicate as PredicateV1; +use schemars::JsonSchema; +use strum_macros::{Display, EnumIter}; + +#[derive(serde::Deserialize, serde::Serialize, Debug, Clone, Copy, JsonSchema, PywrVisitAll, Display, EnumIter)] +pub enum Predicate { + #[serde(alias = "<")] + LT, + #[serde(alias = ">")] + GT, + #[serde(alias = "==")] + EQ, + #[serde(alias = "<=")] + LE, + #[serde(alias = ">=")] + GE, +} + +impl From for Predicate { + fn from(v1: PredicateV1) -> Self { + match v1 { + PredicateV1::LT => Predicate::LT, + PredicateV1::GT => Predicate::GT, + PredicateV1::EQ => Predicate::EQ, + PredicateV1::LE => Predicate::LE, + PredicateV1::GE => Predicate::GE, + } + } +} + +#[cfg(feature = "core")] +impl From for pywr_core::predicate::Predicate { + fn from(p: Predicate) -> Self { + match p { + Predicate::LT => pywr_core::predicate::Predicate::LessThan, + Predicate::GT => pywr_core::predicate::Predicate::GreaterThan, + Predicate::EQ => pywr_core::predicate::Predicate::EqualTo, + Predicate::LE => pywr_core::predicate::Predicate::LessThanOrEqualTo, + Predicate::GE => pywr_core::predicate::Predicate::GreaterThanOrEqualTo, + } + } +} diff --git a/pywr-schema/src/timeseries/align_and_resample.rs b/pywr-schema/src/timeseries/align_and_resample.rs index bc7a7edb..ed3067ef 100644 --- a/pywr-schema/src/timeseries/align_and_resample.rs +++ b/pywr-schema/src/timeseries/align_and_resample.rs @@ -93,7 +93,7 @@ pub fn align_and_resample( fn slice_start(df: DataFrame, time_col: &str, domain: &TimeDomain) -> Result { let start = domain - .first_timestep() + .first() .ok_or_else(|| TimeseriesError::NoTimestepsDefined)? .date; let df = df.clone().lazy().filter(col(time_col).gt_eq(lit(start))).collect()?; @@ -102,7 +102,7 @@ fn slice_start(df: DataFrame, time_col: &str, domain: &TimeDomain) -> Result Result { let end = domain - .last_timestep() + .last() .ok_or_else(|| TimeseriesError::NoTimestepsDefined)? .date; let df = df.clone().lazy().filter(col(time_col).lt_eq(lit(end))).collect()?; diff --git a/pywr-schema/tests/csv2.json b/pywr-schema/tests/csv2.json index 7ad9455c..39a6ba5b 100644 --- a/pywr-schema/tests/csv2.json +++ b/pywr-schema/tests/csv2.json @@ -71,6 +71,7 @@ { "name": "nodes", "aggregator": { + "type": "Periodic", "freq": { "type": "Monthly" }, diff --git a/pywr-schema/tests/csv3.json b/pywr-schema/tests/csv3.json index 8c3edf27..670c9443 100644 --- a/pywr-schema/tests/csv3.json +++ b/pywr-schema/tests/csv3.json @@ -71,6 +71,7 @@ { "name": "nodes-monthly-mean", "aggregator": { + "type": "Periodic", "freq": { "type": "Monthly" }, @@ -88,6 +89,7 @@ { "name": "nodes-annual-mean", "aggregator": { + "type": "Periodic", "freq": { "type": "Annual" }, diff --git a/pywr-schema/tests/memory1.json b/pywr-schema/tests/memory1.json index 306bf4e3..abfba192 100644 --- a/pywr-schema/tests/memory1.json +++ b/pywr-schema/tests/memory1.json @@ -71,6 +71,7 @@ { "name": "nodes", "aggregator": { + "type": "Periodic", "freq": { "type": "Annual" }, @@ -78,6 +79,7 @@ "type": "CountNonZero" }, "child": { + "type": "Periodic", "freq": { "type": "Days", "days": 4 diff --git a/pywr-schema/tests/reservoir-failure-levels1.json b/pywr-schema/tests/reservoir-failure-levels1.json index d2740644..dd2be286 100644 --- a/pywr-schema/tests/reservoir-failure-levels1.json +++ b/pywr-schema/tests/reservoir-failure-levels1.json @@ -96,6 +96,7 @@ { "name": "reservoir-metrics", "aggregator": { + "type": "Periodic", "freq": { "type": "Annual" },