use crate::core::TimeSeries;
use crate::error::{ForecastError, Result};
use crate::models::laplace::LaplaceForecaster;
use crate::models::traits::Forecaster;
use crate::models::DistributionalForecaster;
use std::collections::HashMap;
pub struct GlobalLaplace {
factory: Box<dyn Fn() -> LaplaceForecaster + Send + Sync>,
fitted: HashMap<String, LaplaceForecaster>,
panel_calibration_scale: Option<f64>,
}
impl GlobalLaplace {
pub fn new<F>(factory: F) -> Self
where
F: Fn() -> LaplaceForecaster + Send + Sync + 'static,
{
Self {
factory: Box::new(factory),
fitted: HashMap::new(),
panel_calibration_scale: None,
}
}
pub fn fit_series(&mut self, id: impl Into<String>, series: &TimeSeries) -> Result<()> {
let mut m = (self.factory)();
m.fit(series)?;
self.fitted.insert(id.into(), m);
Ok(())
}
pub fn fit_panel<'a, I>(&mut self, panel: I) -> usize
where
I: IntoIterator<Item = (String, &'a TimeSeries)>,
{
let mut ok = 0;
for (id, ts) in panel {
if self.fit_series(id, ts).is_ok() {
ok += 1;
}
}
ok
}
pub fn predict_series(&self, id: &str, horizon: usize) -> Result<crate::core::Forecast> {
let m = self.fitted.get(id).ok_or_else(|| {
ForecastError::InvalidParameter(format!("no fit found for series id `{id}`"))
})?;
m.predict(horizon)
}
pub fn forecast_dist_series(
&self,
id: &str,
horizon: usize,
) -> Result<Vec<super::dist::GaussianMixture>> {
let m = self.fitted.get(id).ok_or_else(|| {
ForecastError::InvalidParameter(format!("no fit found for series id `{id}`"))
})?;
m.forecast_dist(horizon)
}
pub fn panel_calibration_scale(&self) -> Option<f64> {
self.panel_calibration_scale
}
pub fn n_fitted(&self) -> usize {
self.fitted.len()
}
}
pub struct MetaLearnerScaffold;
impl MetaLearnerScaffold {
pub fn pick_family(
zero_fraction: f64,
_trend_strength: f64,
) -> crate::models::smart::SelectedFamily {
use crate::models::smart::SelectedFamily;
if zero_fraction > 0.4 {
SelectedFamily::IntermittentNegBinomial
} else {
SelectedFamily::RegularNormal
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, TimeZone, Utc};
fn ts_ar1(n: usize, phi: f64) -> TimeSeries {
let base = Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap();
let mut vals = Vec::with_capacity(n);
let mut y = 0.0;
for i in 0..n {
let eps = ((i as f64 * 12.9898).sin() * 43758.5453).fract() - 0.5;
y = phi * y + eps + 5.0;
vals.push(y);
}
let stamps: Vec<_> = (0..n).map(|i| base + Duration::hours(i as i64)).collect();
TimeSeries::univariate(stamps, vals).unwrap()
}
#[test]
fn fits_a_small_panel_and_predicts_per_series() {
let mut g = GlobalLaplace::new(|| LaplaceForecaster::new().auto());
let ts_a = ts_ar1(120, 0.4);
let ts_b = ts_ar1(120, 0.6);
g.fit_series("A", &ts_a).unwrap();
g.fit_series("B", &ts_b).unwrap();
assert_eq!(g.n_fitted(), 2);
let fc = g.predict_series("A", 5).unwrap();
assert_eq!(fc.primary().len(), 5);
}
#[test]
fn scaffold_pick_family_matches_smart_rules() {
use crate::models::smart::SelectedFamily;
assert_eq!(
MetaLearnerScaffold::pick_family(0.6, 0.0),
SelectedFamily::IntermittentNegBinomial
);
assert_eq!(
MetaLearnerScaffold::pick_family(0.1, 0.2),
SelectedFamily::RegularNormal
);
}
}