use super::traits::{Recency, TrendComponent};
use crate::error::{ForecastError, Result};
#[derive(Debug, Clone)]
pub struct ExponentialTrend {
recency: Recency,
a: f64,
b: f64,
fitted: Vec<f64>,
n_train: usize,
r_squared: f64,
}
impl ExponentialTrend {
pub fn new() -> Self {
Self {
recency: Recency::Fraction(0.3),
a: 0.0,
b: 0.0,
fitted: Vec::new(),
n_train: 0,
r_squared: 0.0,
}
}
pub fn with_recency(mut self, recency: Recency) -> Self {
self.recency = recency;
self
}
pub fn growth_rate(&self) -> f64 {
self.b
}
}
impl Default for ExponentialTrend {
fn default() -> Self {
Self::new()
}
}
impl TrendComponent for ExponentialTrend {
fn fit_trend(&mut self, values: &[f64]) -> Result<()> {
if values.is_empty() {
return Err(ForecastError::EmptyData);
}
let n = values.len();
if n == 1 {
if values[0] <= 0.0 {
return Err(ForecastError::InvalidParameter(
"ExponentialTrend requires all positive values in the recency window"
.to_string(),
));
}
self.a = values[0].ln();
self.b = 0.0;
self.fitted = vec![values[0]];
self.n_train = 1;
self.r_squared = 1.0;
return Ok(());
}
let (rec_start, rec_end) = self.recency.resolve_with_data(values);
let window = &values[rec_start..rec_end];
let log_y: Vec<f64> = window
.iter()
.map(|&y| {
if y <= 0.0 {
Err(ForecastError::InvalidParameter(
"ExponentialTrend requires all positive values in the recency window"
.to_string(),
))
} else {
Ok(y.ln())
}
})
.collect::<Result<Vec<f64>>>()?;
let w = log_y.len();
let w_f = w as f64;
let t_sum: f64 = (rec_start..rec_end).map(|i| i as f64).sum();
let t_mean = t_sum / w_f;
let logy_mean = log_y.iter().sum::<f64>() / w_f;
let mut ss_tt = 0.0;
let mut ss_ty = 0.0;
for (j, &ly) in log_y.iter().enumerate() {
let t = (rec_start + j) as f64;
let dt = t - t_mean;
let dy = ly - logy_mean;
ss_tt += dt * dt;
ss_ty += dt * dy;
}
let slope = if ss_tt.abs() < 1e-15 {
0.0
} else {
ss_ty / ss_tt
};
let intercept = logy_mean - slope * t_mean;
self.a = intercept;
self.b = slope;
self.fitted = (0..n)
.map(|i| (intercept + slope * i as f64).exp())
.collect();
self.n_train = n;
let window_fitted = &self.fitted[rec_start..rec_end];
let mean_y = window.iter().sum::<f64>() / w_f;
let ss_tot: f64 = window.iter().map(|&y| (y - mean_y).powi(2)).sum();
let ss_res: f64 = window
.iter()
.zip(window_fitted.iter())
.map(|(&y, &f)| (y - f).powi(2))
.sum();
self.r_squared = if ss_tot < 1e-12 {
if ss_res < 1e-12 {
1.0
} else {
0.0
}
} else {
1.0 - ss_res / ss_tot
};
Ok(())
}
fn fitted_trend(&self) -> &[f64] {
&self.fitted
}
fn predict_trend(&self, n_ahead: usize) -> Vec<f64> {
(0..n_ahead)
.map(|i| (self.a + self.b * (self.n_train + i) as f64).exp())
.collect()
}
fn trend_features(&self) -> Vec<(&str, f64)> {
vec![
("exponential_growth_rate", self.b),
("exponential_r_squared", self.r_squared),
]
}
fn trend_name(&self) -> &str {
"exponential"
}
fn n_params(&self) -> usize {
2
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_abs_diff_eq;
#[test]
fn pure_exponential_coefficients() {
let values: Vec<f64> = (0..100).map(|t| 2.0 * (0.05 * t as f64).exp()).collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
assert_abs_diff_eq!(trend.a, 2.0_f64.ln(), epsilon = 1e-10);
assert_abs_diff_eq!(trend.b, 0.05, epsilon = 1e-10);
assert_abs_diff_eq!(trend.growth_rate(), 0.05, epsilon = 1e-10);
}
#[test]
fn fitted_values_close_to_original() {
let values: Vec<f64> = (0..50).map(|t| 2.0 * (0.05 * t as f64).exp()).collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
let fitted = trend.fitted_trend();
assert_eq!(fitted.len(), 50);
for (i, (&f, &v)) in fitted.iter().zip(values.iter()).enumerate() {
assert_abs_diff_eq!(f, v, epsilon = 1e-8);
let _ = i;
}
}
#[test]
fn predict_extrapolates_correctly() {
let n = 50;
let values: Vec<f64> = (0..n).map(|t| 2.0 * (0.05 * t as f64).exp()).collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
let forecast = trend.predict_trend(5);
assert_eq!(forecast.len(), 5);
for (i, &f) in forecast.iter().enumerate() {
let expected = 2.0 * (0.05 * (n + i) as f64).exp();
assert_abs_diff_eq!(f, expected, epsilon = 1e-6);
}
}
#[test]
fn non_positive_values_error() {
let values = vec![1.0, 2.0, 0.0, 4.0, 5.0];
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
let result = trend.fit_trend(&values);
assert!(
matches!(result, Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("positive"))
);
}
#[test]
fn negative_values_error() {
let values = vec![1.0, -2.0, 3.0];
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
let result = trend.fit_trend(&values);
assert!(matches!(result, Err(ForecastError::InvalidParameter(_))));
}
#[test]
fn recency_window_works() {
let n = 100;
let values: Vec<f64> = (0..n).map(|t| 3.0 * (0.02 * t as f64).exp()).collect();
let mut trend_full = ExponentialTrend::new().with_recency(Recency::Full);
trend_full.fit_trend(&values).unwrap();
let mut trend_recent = ExponentialTrend::new().with_recency(Recency::Fraction(0.3));
trend_recent.fit_trend(&values).unwrap();
assert_abs_diff_eq!(trend_full.growth_rate(), 0.02, epsilon = 1e-10);
assert_abs_diff_eq!(trend_recent.growth_rate(), 0.02, epsilon = 1e-6);
assert_eq!(trend_recent.fitted_trend().len(), n);
}
#[test]
fn recency_window_different_rate() {
let n = 100;
let mut values = Vec::with_capacity(n);
for t in 0..70 {
values.push((0.1 * t as f64).exp());
}
let c = (0.1 * 70.0_f64).exp() / (0.02 * 70.0_f64).exp();
for t in 70..100 {
values.push(c * (0.02 * t as f64).exp());
}
let mut trend = ExponentialTrend::new().with_recency(Recency::Fraction(0.3));
trend.fit_trend(&values).unwrap();
assert_abs_diff_eq!(trend.growth_rate(), 0.02, epsilon = 1e-4);
}
#[test]
fn n_params_returns_2() {
let trend = ExponentialTrend::new();
assert_eq!(trend.n_params(), 2);
}
#[test]
fn features_extraction() {
let values: Vec<f64> = (0..50).map(|t| 2.0 * (0.05 * t as f64).exp()).collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
let features = trend.trend_features();
let map: std::collections::HashMap<&str, f64> = features.into_iter().collect();
assert!(map.contains_key("exponential_growth_rate"));
assert!(map.contains_key("exponential_r_squared"));
assert_abs_diff_eq!(map["exponential_growth_rate"], 0.05, epsilon = 1e-10);
assert!(
map["exponential_r_squared"] > 0.999,
"R^2 should be ~1 for pure exponential data, got {}",
map["exponential_r_squared"]
);
}
#[test]
fn features_before_fit() {
let trend = ExponentialTrend::new();
let features = trend.trend_features();
assert_eq!(features.len(), 2);
let map: std::collections::HashMap<&str, f64> = features.into_iter().collect();
assert_abs_diff_eq!(map["exponential_growth_rate"], 0.0, epsilon = 1e-15);
assert_abs_diff_eq!(map["exponential_r_squared"], 0.0, epsilon = 1e-15);
}
#[test]
fn empty_data_error() {
let mut trend = ExponentialTrend::new();
let result = trend.fit_trend(&[]);
assert!(matches!(result, Err(ForecastError::EmptyData)));
}
#[test]
fn single_point() {
let mut trend = ExponentialTrend::new();
trend.fit_trend(&[5.0]).unwrap();
assert_eq!(trend.fitted_trend().len(), 1);
assert_abs_diff_eq!(trend.fitted_trend()[0], 5.0, epsilon = 1e-10);
assert_abs_diff_eq!(trend.growth_rate(), 0.0, epsilon = 1e-15);
assert_abs_diff_eq!(trend.r_squared, 1.0, epsilon = 1e-15);
let forecast = trend.predict_trend(3);
assert_eq!(forecast.len(), 3);
for &f in &forecast {
assert_abs_diff_eq!(f, 5.0, epsilon = 1e-10);
}
}
#[test]
fn single_non_positive_point_error() {
let mut trend = ExponentialTrend::new();
let result = trend.fit_trend(&[0.0]);
assert!(matches!(result, Err(ForecastError::InvalidParameter(_))));
let mut trend2 = ExponentialTrend::new();
let result2 = trend2.fit_trend(&[-3.0]);
assert!(matches!(result2, Err(ForecastError::InvalidParameter(_))));
}
#[test]
fn trend_name_is_correct() {
let trend = ExponentialTrend::new();
assert_eq!(trend.trend_name(), "exponential");
}
#[test]
fn r_squared_pure_exponential() {
let values: Vec<f64> = (0..50).map(|t| 10.0 * (0.03 * t as f64).exp()).collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
assert_abs_diff_eq!(trend.r_squared, 1.0, epsilon = 1e-6);
}
#[test]
fn r_squared_noisy_data() {
let values: Vec<f64> = (0..100)
.map(|t| {
let base = 2.0 * (0.03 * t as f64).exp();
base + 0.5 * (t as f64 * 1.7).sin()
})
.collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
assert!(
trend.r_squared < 1.0,
"R^2 should be less than 1 for noisy data"
);
assert!(
trend.r_squared > 0.5,
"R^2 should still be reasonable for mostly exponential data"
);
}
#[test]
fn predict_zero_ahead() {
let values: Vec<f64> = (0..20).map(|t| (0.1 * t as f64).exp()).collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
let forecast = trend.predict_trend(0);
assert!(forecast.is_empty());
}
#[test]
fn predict_before_fit() {
let trend = ExponentialTrend::new();
let forecast = trend.predict_trend(5);
assert_eq!(forecast.len(), 5);
for &f in &forecast {
assert_abs_diff_eq!(f, 1.0, epsilon = 1e-15);
}
}
#[test]
fn default_same_as_new() {
let a = ExponentialTrend::new();
let b = ExponentialTrend::default();
assert_eq!(a.recency, b.recency);
assert_abs_diff_eq!(a.a, b.a, epsilon = 1e-15);
assert_abs_diff_eq!(a.b, b.b, epsilon = 1e-15);
}
#[test]
fn exponential_decay() {
let values: Vec<f64> = (0..80).map(|t| 10.0 * (-0.03 * t as f64).exp()).collect();
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
assert_abs_diff_eq!(trend.growth_rate(), -0.03, epsilon = 1e-10);
assert_abs_diff_eq!(trend.r_squared, 1.0, epsilon = 1e-6);
}
#[test]
fn constant_positive_data() {
let values = vec![5.0; 30];
let mut trend = ExponentialTrend::new().with_recency(Recency::Full);
trend.fit_trend(&values).unwrap();
assert_abs_diff_eq!(trend.growth_rate(), 0.0, epsilon = 1e-10);
for &f in trend.fitted_trend() {
assert_abs_diff_eq!(f, 5.0, epsilon = 1e-10);
}
}
}