use super::MLError;
use crate::DataFrame;
pub struct TimeSeriesForecaster {
pub window_size: usize,
pub forecast_horizon: usize,
pub trend_analysis: Option<TrendAnalysis>,
pub seasonality_analysis: Option<SeasonalityAnalysis>,
}
#[derive(Debug, Clone)]
pub struct TrendAnalysis {
pub trend_type: TrendType,
pub strength: f64,
pub direction: TrendDirection,
pub coefficients: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum TrendType {
Linear,
Exponential,
Logarithmic,
None,
}
#[derive(Debug, Clone, PartialEq)]
pub enum TrendDirection {
Increasing,
Decreasing,
Stable,
}
#[derive(Debug, Clone)]
pub struct SeasonalityAnalysis {
pub has_seasonality: bool,
pub period: Option<usize>,
pub strength: f64,
pub pattern: Vec<f64>,
}
#[derive(Debug, Clone)]
pub struct TimeSeriesDecomposition {
pub trend: Vec<f64>,
pub seasonal: Vec<f64>,
pub residual: Vec<f64>,
pub original: Vec<f64>,
}
impl TimeSeriesForecaster {
pub fn new(window_size: usize, forecast_horizon: usize) -> Self {
Self {
window_size,
forecast_horizon,
trend_analysis: None,
seasonality_analysis: None,
}
}
pub fn analyze(&mut self, data: &DataFrame) -> Result<(), MLError> {
if data.height() < self.window_size {
return Err(MLError::InsufficientData(
format!(
"Need at least {} data points for analysis",
self.window_size
),
));
}
let values = self.extract_values(data)?;
self.trend_analysis = Some(self.analyze_trend(&values)?);
self.seasonality_analysis = Some(self.analyze_seasonality(&values)?);
Ok(())
}
pub fn forecast(&self, data: &DataFrame) -> Result<Vec<f64>, MLError> {
if self.trend_analysis.is_none() {
return Err(MLError::ModelNotTrained);
}
let values = self.extract_values(data)?;
let mut forecasts = Vec::with_capacity(self.forecast_horizon);
for i in 0..self.forecast_horizon {
let forecast = self.generate_single_forecast(&values, i)?;
forecasts.push(forecast);
}
Ok(forecasts)
}
pub fn decompose(&self, data: &DataFrame) -> Result<TimeSeriesDecomposition, MLError> {
let values = self.extract_values(data)?;
let _n = values.len();
let trend = self.calculate_trend_component(&values);
let seasonal = self.calculate_seasonal_component(&values);
let residual = self.calculate_residual_component(&values, &trend, &seasonal);
Ok(TimeSeriesDecomposition {
trend,
seasonal,
residual,
original: values,
})
}
fn extract_values(&self, data: &DataFrame) -> Result<Vec<f64>, MLError> {
let columns = data.get_columns();
for column in columns {
if let Ok(series) = column.f64() {
return Ok(series.into_iter().filter_map(|v| v).collect());
}
}
Err(MLError::InvalidData(
"No numeric column found for time series analysis".to_string(),
))
}
fn analyze_trend(&self, values: &[f64]) -> Result<TrendAnalysis, MLError> {
if values.len() < 2 {
return Err(MLError::InsufficientData(
"Need at least 2 points for trend analysis".to_string(),
));
}
let (slope, intercept, r_squared) = self.calculate_linear_regression(values);
let trend_type = if r_squared > 0.7 {
TrendType::Linear
} else {
TrendType::None
};
let direction = if slope > 0.01 {
TrendDirection::Increasing
} else if slope < -0.01 {
TrendDirection::Decreasing
} else {
TrendDirection::Stable
};
Ok(TrendAnalysis {
trend_type,
strength: r_squared,
direction,
coefficients: vec![intercept, slope],
})
}
fn analyze_seasonality(&self, values: &[f64]) -> Result<SeasonalityAnalysis, MLError> {
let n = values.len();
if n < 12 {
return Ok(SeasonalityAnalysis {
has_seasonality: false,
period: None,
strength: 0.0,
pattern: Vec::new(),
});
}
let max_period = (n / 4).min(24);
let mut best_period = None;
let mut best_strength = 0.0;
for period in 2..=max_period {
let strength = self.calculate_seasonal_strength(values, period);
if strength > best_strength {
best_strength = strength;
best_period = Some(period);
}
}
let has_seasonality = best_strength > 0.3;
let pattern = if has_seasonality {
self.extract_seasonal_pattern(values, best_period.unwrap())
} else {
Vec::new()
};
Ok(SeasonalityAnalysis {
has_seasonality,
period: best_period,
strength: best_strength,
pattern,
})
}
fn calculate_linear_regression(&self, values: &[f64]) -> (f64, f64, f64) {
let n = values.len() as f64;
let x_values: Vec<f64> = (0..values.len()).map(|i| i as f64).collect();
let sum_x: f64 = x_values.iter().sum();
let sum_y: f64 = values.iter().sum();
let sum_xy: f64 = x_values.iter().zip(values.iter()).map(|(a, b)| a * b).sum();
let sum_x2: f64 = x_values.iter().map(|a| a * a).sum();
let _sum_y2: f64 = values.iter().map(|a| a * a).sum();
let slope = (n * sum_xy - sum_x * sum_y) / (n * sum_x2 - sum_x * sum_x);
let intercept = (sum_y - slope * sum_x) / n;
let y_mean = sum_y / n;
let ss_tot: f64 = values.iter().map(|yi| (yi - y_mean).powi(2)).sum();
let ss_res: f64 = x_values
.iter()
.zip(values.iter())
.map(|(xi, yi)| (yi - (slope * xi + intercept)).powi(2))
.sum();
let r_squared = 1.0 - (ss_res / ss_tot);
(slope, intercept, r_squared)
}
fn calculate_seasonal_strength(&self, values: &[f64], period: usize) -> f64 {
let n = values.len();
if n < period * 2 {
return 0.0;
}
let mut autocorr_sum = 0.0;
let mut count = 0;
for i in 0..(n - period) {
autocorr_sum += values[i] * values[i + period];
count += 1;
}
if count == 0 {
return 0.0;
}
autocorr_sum / count as f64
}
fn extract_seasonal_pattern(&self, values: &[f64], period: usize) -> Vec<f64> {
let mut pattern = vec![0.0; period];
let mut counts = vec![0; period];
for (i, &value) in values.iter().enumerate() {
let phase = i % period;
pattern[phase] += value;
counts[phase] += 1;
}
for i in 0..period {
if counts[i] > 0 {
pattern[i] /= counts[i] as f64;
}
}
pattern
}
fn generate_single_forecast(&self, values: &[f64], step: usize) -> Result<f64, MLError> {
let trend = &self.trend_analysis.as_ref().unwrap();
let seasonal = &self.seasonality_analysis.as_ref().unwrap();
let n = values.len();
let trend_forecast = trend.coefficients[0] + trend.coefficients[1] * (n + step) as f64;
let seasonal_component = if seasonal.has_seasonality {
let period = seasonal.period.unwrap_or(12);
seasonal.pattern[(n + step) % period]
} else {
0.0
};
Ok(trend_forecast + seasonal_component)
}
fn calculate_trend_component(&self, values: &[f64]) -> Vec<f64> {
let window_size = (values.len() / 10).max(3).min(20);
let mut trend = Vec::new();
for i in 0..values.len() {
let start = i.saturating_sub(window_size / 2);
let end = (i + window_size / 2 + 1).min(values.len());
let window_sum: f64 = values[start..end].iter().sum();
let window_avg = window_sum / (end - start) as f64;
trend.push(window_avg);
}
trend
}
fn calculate_seasonal_component(&self, values: &[f64]) -> Vec<f64> {
let seasonal = &self.seasonality_analysis.as_ref().unwrap();
let mut seasonal_component = vec![0.0; values.len()];
if seasonal.has_seasonality {
let period = seasonal.period.unwrap_or(12);
for i in 0..values.len() {
let phase = i % period;
seasonal_component[i] = seasonal.pattern[phase];
}
}
seasonal_component
}
fn calculate_residual_component(
&self,
values: &[f64],
trend: &[f64],
seasonal: &[f64],
) -> Vec<f64> {
values
.iter()
.zip(trend.iter())
.zip(seasonal.iter())
.map(|((&original, &trend_val), &seasonal_val)| {
original - trend_val - seasonal_val
})
.collect()
}
}