use chrono::{DateTime, Datelike, Utc};
use std::collections::HashMap;
use std::f64::consts::PI;
use crate::core::{CalendarAnnotations, TimeSeries};
#[derive(Debug, Clone)]
pub struct FeatureGenerator {
specs: Vec<FeatureSpec>,
}
#[derive(Debug, Clone)]
enum FeatureSpec {
Fourier {
period: usize,
order: usize,
},
DayOfWeek,
MonthOfYear,
Quarter,
Holiday {
dates: Vec<DateTime<Utc>>,
name: String,
},
Cyclical(TimeComponent),
Binary(BinaryIndicator),
Advanced(AdvancedFeature),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum TimeComponent {
Month,
Quarter,
Semester,
WeekOfYear,
DayOfWeek,
DayOfMonth,
DayOfYear,
Hour,
Minute,
Second,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum BinaryIndicator {
MonthStart,
MonthEnd,
QuarterStart,
QuarterEnd,
YearStart,
YearEnd,
Weekend,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum AdvancedFeature {
LeapYear,
DaysInMonth,
}
impl FeatureGenerator {
pub fn new() -> Self {
Self { specs: Vec::new() }
}
pub fn fourier(mut self, period: usize, order: usize) -> Self {
self.specs.push(FeatureSpec::Fourier { period, order });
self
}
pub fn day_of_week(mut self) -> Self {
self.specs.push(FeatureSpec::DayOfWeek);
self
}
pub fn month_of_year(mut self) -> Self {
self.specs.push(FeatureSpec::MonthOfYear);
self
}
pub fn quarter(mut self) -> Self {
self.specs.push(FeatureSpec::Quarter);
self
}
pub fn cyclical(mut self, component: TimeComponent) -> Self {
self.specs.push(FeatureSpec::Cyclical(component));
self
}
pub fn binary(mut self, indicator: BinaryIndicator) -> Self {
self.specs.push(FeatureSpec::Binary(indicator));
self
}
pub fn advanced(mut self, feature: AdvancedFeature) -> Self {
self.specs.push(FeatureSpec::Advanced(feature));
self
}
pub fn holiday(mut self, name: impl Into<String>, dates: Vec<DateTime<Utc>>) -> Self {
self.specs.push(FeatureSpec::Holiday {
dates,
name: name.into(),
});
self
}
pub fn generate(&self, timestamps: &[DateTime<Utc>]) -> HashMap<String, Vec<f64>> {
let mut result = HashMap::new();
let n = timestamps.len();
for spec in &self.specs {
match spec {
FeatureSpec::Fourier { period, order } => {
let period_f = *period as f64;
for k in 1..=*order {
let freq = 2.0 * PI * k as f64 / period_f;
let mut sin_col = Vec::with_capacity(n);
let mut cos_col = Vec::with_capacity(n);
for (i, _) in timestamps.iter().enumerate() {
let angle = freq * i as f64;
sin_col.push(angle.sin());
cos_col.push(angle.cos());
}
result.insert(format!("fourier_{}_sin{}", period, k), sin_col);
result.insert(format!("fourier_{}_cos{}", period, k), cos_col);
}
}
FeatureSpec::DayOfWeek => {
let names = ["mon", "tue", "wed", "thu", "fri", "sat"];
for (dow, name) in names.iter().enumerate() {
let col: Vec<f64> = timestamps
.iter()
.map(|ts| {
if ts.weekday().num_days_from_monday() as usize == dow {
1.0
} else {
0.0
}
})
.collect();
result.insert(format!("dow_{}", name), col);
}
}
FeatureSpec::MonthOfYear => {
let names = [
"feb", "mar", "apr", "may", "jun", "jul", "aug", "sep", "oct", "nov", "dec",
];
for (i, name) in names.iter().enumerate() {
let month = i + 2; let col: Vec<f64> = timestamps
.iter()
.map(|ts| {
if ts.month() as usize == month {
1.0
} else {
0.0
}
})
.collect();
result.insert(format!("month_{}", name), col);
}
}
FeatureSpec::Quarter => {
for q in 2..=4u32 {
let col: Vec<f64> = timestamps
.iter()
.map(|ts| {
let m = ts.month();
let quarter = (m - 1) / 3 + 1;
if quarter == q {
1.0
} else {
0.0
}
})
.collect();
result.insert(format!("quarter_{}", q), col);
}
}
FeatureSpec::Holiday { dates, name } => {
let holiday_dates: std::collections::HashSet<_> =
dates.iter().map(|d| d.date_naive()).collect();
let col: Vec<f64> = timestamps
.iter()
.map(|ts| {
if holiday_dates.contains(&ts.date_naive()) {
1.0
} else {
0.0
}
})
.collect();
result.insert(format!("holiday_{}", name), col);
}
FeatureSpec::Cyclical(component) => {
use chrono::{Datelike, Timelike};
let (name, period) = match component {
TimeComponent::Month => ("month", 12.0),
TimeComponent::Quarter => ("quarter", 4.0),
TimeComponent::Semester => ("semester", 2.0),
TimeComponent::WeekOfYear => ("week_of_year", 53.0),
TimeComponent::DayOfWeek => ("day_of_week", 7.0),
TimeComponent::DayOfMonth => ("day_of_month", 31.0),
TimeComponent::DayOfYear => ("day_of_year", 366.0),
TimeComponent::Hour => ("hour", 24.0),
TimeComponent::Minute => ("minute", 60.0),
TimeComponent::Second => ("second", 60.0),
};
let mut sin_col = Vec::with_capacity(n);
let mut cos_col = Vec::with_capacity(n);
for ts in timestamps {
let value = match component {
TimeComponent::Month => ts.month() as f64,
TimeComponent::Quarter => ((ts.month() - 1) / 3 + 1) as f64,
TimeComponent::Semester => ((ts.month() - 1) / 6 + 1) as f64,
TimeComponent::WeekOfYear => ts.iso_week().week() as f64,
TimeComponent::DayOfWeek => ts.weekday().num_days_from_monday() as f64,
TimeComponent::DayOfMonth => ts.day() as f64,
TimeComponent::DayOfYear => ts.ordinal() as f64,
TimeComponent::Hour => ts.hour() as f64,
TimeComponent::Minute => ts.minute() as f64,
TimeComponent::Second => ts.second() as f64,
};
let angle = 2.0 * PI * value / period;
sin_col.push(angle.sin());
cos_col.push(angle.cos());
}
result.insert(format!("{}_sin", name), sin_col);
result.insert(format!("{}_cos", name), cos_col);
}
FeatureSpec::Binary(indicator) => {
use chrono::Datelike;
let (name, test_fn): (&str, Box<dyn Fn(&DateTime<Utc>) -> bool>) =
match indicator {
BinaryIndicator::MonthStart => {
("month_start", Box::new(|ts| ts.day() == 1))
}
BinaryIndicator::MonthEnd => (
"month_end",
Box::new(|ts| {
let max_day = crate::core::time_series::days_in_month_pub(
ts.year(),
ts.month(),
);
ts.day() == max_day
}),
),
BinaryIndicator::QuarterStart => (
"quarter_start",
Box::new(|ts| {
ts.day() == 1 && matches!(ts.month(), 1 | 4 | 7 | 10)
}),
),
BinaryIndicator::QuarterEnd => (
"quarter_end",
Box::new(|ts| {
let m = ts.month();
let max_day =
crate::core::time_series::days_in_month_pub(ts.year(), m);
ts.day() == max_day && matches!(m, 3 | 6 | 9 | 12)
}),
),
BinaryIndicator::YearStart => (
"year_start",
Box::new(|ts| ts.month() == 1 && ts.day() == 1),
),
BinaryIndicator::YearEnd => (
"year_end",
Box::new(|ts| ts.month() == 12 && ts.day() == 31),
),
BinaryIndicator::Weekend => (
"weekend",
Box::new(|ts| {
matches!(
ts.weekday(),
chrono::Weekday::Sat | chrono::Weekday::Sun
)
}),
),
};
let col: Vec<f64> = timestamps
.iter()
.map(|ts| if test_fn(ts) { 1.0 } else { 0.0 })
.collect();
result.insert(name.to_string(), col);
}
FeatureSpec::Advanced(feature) => {
use chrono::Datelike;
let (name, compute_fn): (&str, Box<dyn Fn(&DateTime<Utc>) -> f64>) =
match feature {
AdvancedFeature::LeapYear => (
"leap_year",
Box::new(|ts| {
if crate::core::time_series::is_leap_year_pub(ts.year()) {
1.0
} else {
0.0
}
}),
),
AdvancedFeature::DaysInMonth => (
"days_in_month",
Box::new(|ts| {
crate::core::time_series::days_in_month_pub(
ts.year(),
ts.month(),
) as f64
}),
),
};
let col: Vec<f64> = timestamps.iter().map(&*compute_fn).collect();
result.insert(name.to_string(), col);
}
}
}
result
}
pub fn add_to(&self, ts: &mut TimeSeries) {
let features = self.generate(ts.timestamps());
let mut cal = ts
.calendar()
.cloned()
.unwrap_or_else(CalendarAnnotations::new);
for (name, values) in features {
cal = cal.with_regressor(name, values);
}
ts.set_calendar(cal);
}
pub fn feature_names(&self) -> Vec<String> {
let dummy = vec![Utc::now()];
let map = self.generate(&dummy);
let mut names: Vec<String> = map.into_keys().collect();
names.sort();
names
}
}
impl Default for FeatureGenerator {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, TimeZone};
fn daily_timestamps(n: usize) -> Vec<DateTime<Utc>> {
(0..n)
.map(|i| Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap() + Duration::days(i as i64))
.collect()
}
fn monthly_timestamps(n: usize) -> Vec<DateTime<Utc>> {
let mut ts = Vec::with_capacity(n);
let mut year = 2020i32;
let mut month = 1u32;
for _ in 0..n {
ts.push(Utc.with_ymd_and_hms(year, month, 1, 0, 0, 0).unwrap());
month += 1;
if month > 12 {
month = 1;
year += 1;
}
}
ts
}
#[test]
fn fourier_column_count() {
let gen = FeatureGenerator::new().fourier(7, 3).fourier(365, 2);
let feats = gen.generate(&daily_timestamps(30));
assert_eq!(feats.len(), 10);
assert!(feats.contains_key("fourier_7_sin1"));
assert!(feats.contains_key("fourier_7_cos3"));
assert!(feats.contains_key("fourier_365_sin2"));
}
#[test]
fn fourier_values_periodic() {
let gen = FeatureGenerator::new().fourier(7, 1);
let ts = daily_timestamps(14);
let feats = gen.generate(&ts);
let sin1 = &feats["fourier_7_sin1"];
for i in 0..7 {
assert!(
(sin1[i] - sin1[i + 7]).abs() < 1e-10,
"sin1[{}]={} != sin1[{}]={}",
i,
sin1[i],
i + 7,
sin1[i + 7]
);
}
}
#[test]
fn fourier_sin_at_zero_is_zero() {
let gen = FeatureGenerator::new().fourier(12, 1);
let ts = monthly_timestamps(12);
let feats = gen.generate(&ts);
assert!((feats["fourier_12_sin1"][0]).abs() < 1e-10);
assert!((feats["fourier_12_cos1"][0] - 1.0).abs() < 1e-10);
}
#[test]
fn day_of_week_columns() {
let gen = FeatureGenerator::new().day_of_week();
let feats = gen.generate(&daily_timestamps(7));
assert_eq!(feats.len(), 6); assert!(feats.contains_key("dow_mon"));
assert!(feats.contains_key("dow_sat"));
assert!(!feats.contains_key("dow_sun"));
}
#[test]
fn day_of_week_one_hot() {
let gen = FeatureGenerator::new().day_of_week();
let ts = daily_timestamps(7);
let feats = gen.generate(&ts);
assert_eq!(feats["dow_mon"][0], 1.0);
assert_eq!(feats["dow_tue"][0], 0.0);
assert_eq!(feats["dow_tue"][1], 1.0); assert_eq!(feats["dow_sat"][5], 1.0);
for col in feats.values() {
assert_eq!(col[6], 0.0);
}
}
#[test]
fn month_of_year_columns() {
let gen = FeatureGenerator::new().month_of_year();
let feats = gen.generate(&monthly_timestamps(12));
assert_eq!(feats.len(), 11); assert!(feats.contains_key("month_feb"));
assert!(feats.contains_key("month_dec"));
assert!(!feats.contains_key("month_jan"));
}
#[test]
fn month_of_year_one_hot() {
let gen = FeatureGenerator::new().month_of_year();
let ts = monthly_timestamps(12);
let feats = gen.generate(&ts);
for col in feats.values() {
assert_eq!(col[0], 0.0);
}
assert_eq!(feats["month_feb"][1], 1.0);
assert_eq!(feats["month_mar"][1], 0.0);
assert_eq!(feats["month_dec"][11], 1.0);
}
#[test]
fn quarter_columns() {
let gen = FeatureGenerator::new().quarter();
let feats = gen.generate(&monthly_timestamps(12));
assert_eq!(feats.len(), 3); }
#[test]
fn quarter_one_hot() {
let gen = FeatureGenerator::new().quarter();
let ts = monthly_timestamps(12);
let feats = gen.generate(&ts);
for i in 0..3 {
assert_eq!(feats["quarter_2"][i], 0.0);
assert_eq!(feats["quarter_3"][i], 0.0);
assert_eq!(feats["quarter_4"][i], 0.0);
}
for i in 3..6 {
assert_eq!(feats["quarter_2"][i], 1.0);
}
for i in 9..12 {
assert_eq!(feats["quarter_4"][i], 1.0);
}
}
#[test]
fn holiday_indicator() {
let xmas = Utc.with_ymd_and_hms(2024, 1, 3, 0, 0, 0).unwrap();
let gen = FeatureGenerator::new().holiday("xmas", vec![xmas]);
let ts = daily_timestamps(7);
let feats = gen.generate(&ts);
let col = &feats["holiday_xmas"];
assert_eq!(col[2], 1.0); assert_eq!(col[0], 0.0);
assert_eq!(col[4], 0.0);
}
#[test]
fn add_to_attaches_regressors() {
let ts_vec = daily_timestamps(14);
let values: Vec<f64> = (0..14).map(|i| i as f64).collect();
let mut ts = TimeSeries::univariate(ts_vec, values).unwrap();
let gen = FeatureGenerator::new().fourier(7, 2).day_of_week();
gen.add_to(&mut ts);
let regs = ts.all_regressors();
assert_eq!(regs.len(), 10);
assert!(regs.contains_key("fourier_7_sin1"));
assert!(regs.contains_key("dow_fri"));
}
#[test]
fn add_to_preserves_existing_regressors() {
let ts_vec = daily_timestamps(14);
let values: Vec<f64> = (0..14).map(|i| i as f64).collect();
let mut ts = TimeSeries::univariate(ts_vec, values).unwrap();
let cal =
CalendarAnnotations::new().with_regressor("temperature".to_string(), vec![20.0; 14]);
ts.set_calendar(cal);
let gen = FeatureGenerator::new().fourier(7, 1);
gen.add_to(&mut ts);
let regs = ts.all_regressors();
assert!(regs.contains_key("temperature")); assert!(regs.contains_key("fourier_7_sin1")); assert_eq!(regs.len(), 3); }
#[test]
fn feature_names_sorted() {
let gen = FeatureGenerator::new()
.fourier(7, 2)
.day_of_week()
.month_of_year()
.quarter();
let names = gen.feature_names();
let mut sorted = names.clone();
sorted.sort();
assert_eq!(names, sorted);
assert_eq!(names.len(), 24);
}
#[test]
fn generate_for_future_timestamps() {
let gen = FeatureGenerator::new().fourier(7, 1).day_of_week();
let train = daily_timestamps(28);
let future = daily_timestamps(35)[28..].to_vec();
let train_feats = gen.generate(&train);
let future_feats = gen.generate(&future);
assert_eq!(train_feats.len(), future_feats.len());
for key in train_feats.keys() {
assert!(future_feats.contains_key(key), "missing key: {}", key);
}
for col in future_feats.values() {
assert_eq!(col.len(), 7);
}
}
#[test]
fn empty_generator() {
let gen = FeatureGenerator::new();
let feats = gen.generate(&daily_timestamps(10));
assert!(feats.is_empty());
assert!(gen.feature_names().is_empty());
}
use crate::utils::ols::ols_fit;
fn assert_ols_recovers(
features: &HashMap<String, Vec<f64>>,
true_intercept: f64,
true_coeffs: &HashMap<String, f64>,
tol: f64,
) {
let n = features.values().next().unwrap().len();
let mut y = vec![true_intercept; n];
for (name, coeff) in true_coeffs {
let col = &features[name];
for (i, yi) in y.iter_mut().enumerate() {
*yi += coeff * col[i];
}
}
let result = ols_fit(&y, features).unwrap();
assert!(
(result.intercept - true_intercept).abs() < tol,
"intercept: expected {}, got {} (tol {})",
true_intercept,
result.intercept,
tol,
);
for (name, &expected) in true_coeffs {
let idx = result
.regressor_names
.iter()
.position(|n| n == name)
.unwrap_or_else(|| panic!("missing regressor '{}'", name));
assert!(
(result.coefficients[idx] - expected).abs() < tol,
"coeff '{}': expected {}, got {} (tol {})",
name,
expected,
result.coefficients[idx],
tol,
);
}
}
#[test]
fn ols_recovers_fourier_effects() {
let gen = FeatureGenerator::new().fourier(12, 1);
let ts = monthly_timestamps(120); let features = gen.generate(&ts);
let mut true_coeffs = HashMap::new();
true_coeffs.insert("fourier_12_sin1".to_string(), 5.0);
true_coeffs.insert("fourier_12_cos1".to_string(), -3.0);
assert_ols_recovers(&features, 100.0, &true_coeffs, 0.01);
}
#[test]
fn ols_recovers_day_of_week_effects() {
let gen = FeatureGenerator::new().day_of_week();
let ts = daily_timestamps(364); let features = gen.generate(&ts);
let mut true_coeffs = HashMap::new();
true_coeffs.insert("dow_mon".to_string(), 10.0);
true_coeffs.insert("dow_tue".to_string(), 8.0);
true_coeffs.insert("dow_wed".to_string(), 6.0);
true_coeffs.insert("dow_thu".to_string(), 4.0);
true_coeffs.insert("dow_fri".to_string(), 12.0);
true_coeffs.insert("dow_sat".to_string(), -5.0);
assert_ols_recovers(&features, 50.0, &true_coeffs, 0.01);
}
#[test]
fn ols_recovers_month_effects() {
let gen = FeatureGenerator::new().month_of_year();
let ts = monthly_timestamps(120); let features = gen.generate(&ts);
let mut true_coeffs = HashMap::new();
let month_effects = [
("month_feb", 2.0),
("month_mar", 5.0),
("month_apr", 10.0),
("month_may", 15.0),
("month_jun", 18.0),
("month_jul", 20.0),
("month_aug", 19.0),
("month_sep", 14.0),
("month_oct", 8.0),
("month_nov", 3.0),
("month_dec", -1.0),
];
for (name, effect) in &month_effects {
true_coeffs.insert(name.to_string(), *effect);
}
assert_ols_recovers(&features, 200.0, &true_coeffs, 0.01);
}
#[test]
fn ols_recovers_quarter_effects() {
let gen = FeatureGenerator::new().quarter();
let ts = monthly_timestamps(48); let features = gen.generate(&ts);
let mut true_coeffs = HashMap::new();
true_coeffs.insert("quarter_2".to_string(), 15.0);
true_coeffs.insert("quarter_3".to_string(), 25.0);
true_coeffs.insert("quarter_4".to_string(), 10.0);
assert_ols_recovers(&features, 80.0, &true_coeffs, 0.01);
}
#[test]
fn ols_recovers_holiday_effect() {
let ts = daily_timestamps(365);
let holiday_dates: Vec<DateTime<Utc>> = [10, 50, 100, 200, 300]
.iter()
.map(|&d| Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap() + Duration::days(d))
.collect();
let gen = FeatureGenerator::new().holiday("promo", holiday_dates);
let features = gen.generate(&ts);
let mut true_coeffs = HashMap::new();
true_coeffs.insert("holiday_promo".to_string(), 50.0);
assert_ols_recovers(&features, 30.0, &true_coeffs, 0.01);
}
#[test]
fn ols_recovers_combined_effects() {
let n = 364; let ts = daily_timestamps(n);
let holiday_dates: Vec<DateTime<Utc>> = (0..52)
.map(|w| Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap() + Duration::days(w * 7))
.collect();
let gen = FeatureGenerator::new()
.fourier(7, 1)
.day_of_week()
.holiday("weekly_event", holiday_dates);
let features = gen.generate(&ts);
let mut true_coeffs = HashMap::new();
true_coeffs.insert("fourier_7_sin1".to_string(), 3.0);
true_coeffs.insert("fourier_7_cos1".to_string(), -2.0);
true_coeffs.insert("dow_mon".to_string(), 5.0);
true_coeffs.insert("dow_tue".to_string(), 4.0);
true_coeffs.insert("dow_wed".to_string(), 3.0);
true_coeffs.insert("dow_thu".to_string(), 2.0);
true_coeffs.insert("dow_fri".to_string(), 6.0);
true_coeffs.insert("dow_sat".to_string(), -1.0);
true_coeffs.insert("holiday_weekly_event".to_string(), 20.0);
let n = features.values().next().unwrap().len();
let mut y = vec![100.0_f64; n];
for (name, coeff) in &true_coeffs {
let col = &features[name];
for (i, yi) in y.iter_mut().enumerate() {
*yi += coeff * col[i];
}
}
let result = ols_fit(&y, &features).unwrap();
let y_hat = result.predict(&features).unwrap();
let max_error = y
.iter()
.zip(y_hat.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f64, f64::max);
assert!(
max_error < 0.01,
"max prediction error: {} (should be < 0.01)",
max_error,
);
}
#[test]
fn ols_recovers_fourier_with_external_regressor() {
let gen = FeatureGenerator::new().fourier(12, 1);
let ts = monthly_timestamps(120);
let mut features = gen.generate(&ts);
let temperature: Vec<f64> = (0..120)
.map(|i| 15.0 + 10.0 * (2.0 * PI * i as f64 / 12.0).cos())
.collect();
features.insert("temperature".to_string(), temperature);
let mut true_coeffs = HashMap::new();
true_coeffs.insert("fourier_12_sin1".to_string(), 4.0);
true_coeffs.insert("fourier_12_cos1".to_string(), 2.0);
true_coeffs.insert("temperature".to_string(), 7.0);
let n = features.values().next().unwrap().len();
let mut y = vec![50.0_f64; n];
for (name, coeff) in &true_coeffs {
let col = &features[name];
for (i, yi) in y.iter_mut().enumerate() {
*yi += coeff * col[i];
}
}
let result = ols_fit(&y, &features).unwrap();
let y_hat = result.predict(&features).unwrap();
let max_error = y
.iter()
.zip(y_hat.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f64, f64::max);
assert!(
max_error < 0.01,
"max prediction error: {} (should be < 0.01)",
max_error,
);
}
#[test]
fn ols_predict_future_with_features() {
let gen = FeatureGenerator::new().fourier(7, 1);
let train_ts = daily_timestamps(70);
let train_features = gen.generate(&train_ts);
let n = 70;
let mut y = vec![100.0; n];
let sin_col = &train_features["fourier_7_sin1"];
let cos_col = &train_features["fourier_7_cos1"];
for i in 0..n {
y[i] += 5.0 * sin_col[i] + 3.0 * cos_col[i];
}
let result = ols_fit(&y, &train_features).unwrap();
let future_ts: Vec<_> = (70..77)
.map(|i| Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap() + Duration::days(i as i64))
.collect();
let future_features = gen.generate(&future_ts);
let predictions = result.predict(&future_features).unwrap();
for (i, &pred) in predictions.iter().enumerate() {
let idx = 70 + i;
let angle = 2.0 * PI / 7.0 * idx as f64;
let expected = 100.0 + 5.0 * angle.sin() + 3.0 * angle.cos();
assert!(
(pred - expected).abs() < 0.01,
"future[{}]: expected {}, got {}",
i,
expected,
pred,
);
}
}
#[test]
fn cyclical_month_january_values() {
let ts = vec![Utc.with_ymd_and_hms(2024, 1, 15, 0, 0, 0).unwrap()];
let gen = FeatureGenerator::new().cyclical(TimeComponent::Month);
let feats = gen.generate(&ts);
let expected_angle = 2.0 * PI * 1.0 / 12.0;
let sin_val = feats["month_sin"][0];
let cos_val = feats["month_cos"][0];
assert!(
(sin_val - expected_angle.sin()).abs() < 1e-10,
"Jan sin: expected {}, got {}",
expected_angle.sin(),
sin_val,
);
assert!(
(cos_val - expected_angle.cos()).abs() < 1e-10,
"Jan cos: expected {}, got {}",
expected_angle.cos(),
cos_val,
);
}
#[test]
fn cyclical_month_december_wraps_near_january() {
let jan = vec![Utc.with_ymd_and_hms(2024, 1, 15, 0, 0, 0).unwrap()];
let dec = vec![Utc.with_ymd_and_hms(2024, 12, 15, 0, 0, 0).unwrap()];
let jul = vec![Utc.with_ymd_and_hms(2024, 7, 15, 0, 0, 0).unwrap()];
let gen = FeatureGenerator::new().cyclical(TimeComponent::Month);
let jan_f = gen.generate(&jan);
let dec_f = gen.generate(&dec);
let jul_f = gen.generate(&jul);
let dist_dec_jan = ((dec_f["month_sin"][0] - jan_f["month_sin"][0]).powi(2)
+ (dec_f["month_cos"][0] - jan_f["month_cos"][0]).powi(2))
.sqrt();
let dist_jul_jan = ((jul_f["month_sin"][0] - jan_f["month_sin"][0]).powi(2)
+ (jul_f["month_cos"][0] - jan_f["month_cos"][0]).powi(2))
.sqrt();
assert!(
dist_dec_jan < dist_jul_jan,
"Dec should be closer to Jan than Jul in cyclical space: dec-jan={}, jul-jan={}",
dist_dec_jan,
dist_jul_jan,
);
}
#[test]
fn cyclical_day_of_week_monday_differs_from_sunday() {
let monday = vec![Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap()];
let sunday = vec![Utc.with_ymd_and_hms(2024, 1, 7, 0, 0, 0).unwrap()];
let gen = FeatureGenerator::new().cyclical(TimeComponent::DayOfWeek);
let mon_f = gen.generate(&monday);
let sun_f = gen.generate(&sunday);
assert!((mon_f["day_of_week_sin"][0] - 0.0).abs() < 1e-10);
assert!((mon_f["day_of_week_cos"][0] - 1.0).abs() < 1e-10);
let expected_angle = 2.0 * PI * 6.0 / 7.0;
assert!(
(sun_f["day_of_week_sin"][0] - expected_angle.sin()).abs() < 1e-10,
"Sunday sin: expected {}, got {}",
expected_angle.sin(),
sun_f["day_of_week_sin"][0],
);
assert!(
(sun_f["day_of_week_cos"][0] - expected_angle.cos()).abs() < 1e-10,
"Sunday cos: expected {}, got {}",
expected_angle.cos(),
sun_f["day_of_week_cos"][0],
);
assert!(
(mon_f["day_of_week_sin"][0] - sun_f["day_of_week_sin"][0]).abs() > 1e-5
|| (mon_f["day_of_week_cos"][0] - sun_f["day_of_week_cos"][0]).abs() > 1e-5,
"Monday and Sunday should have different cyclical encodings"
);
}
#[test]
fn cyclical_hour_zero_equals_hour_24_wrap() {
let hour0 = vec![Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap()];
let gen = FeatureGenerator::new().cyclical(TimeComponent::Hour);
let feats = gen.generate(&hour0);
let sin_0 = feats["hour_sin"][0];
let cos_0 = feats["hour_cos"][0];
assert!(
sin_0.abs() < 1e-10,
"hour=0 sin should be 0.0, got {}",
sin_0
);
assert!(
(cos_0 - 1.0).abs() < 1e-10,
"hour=0 cos should be 1.0, got {}",
cos_0
);
}
#[test]
fn cyclical_all_values_in_range() {
let hourly_ts: Vec<DateTime<Utc>> =
(0..8760) .map(|i| {
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap() + Duration::hours(i as i64)
})
.collect();
let gen = FeatureGenerator::new()
.cyclical(TimeComponent::Month)
.cyclical(TimeComponent::Quarter)
.cyclical(TimeComponent::Semester)
.cyclical(TimeComponent::WeekOfYear)
.cyclical(TimeComponent::DayOfWeek)
.cyclical(TimeComponent::DayOfMonth)
.cyclical(TimeComponent::DayOfYear)
.cyclical(TimeComponent::Hour)
.cyclical(TimeComponent::Minute)
.cyclical(TimeComponent::Second);
let feats = gen.generate(&hourly_ts);
for (name, values) in &feats {
for (i, &v) in values.iter().enumerate() {
assert!(
(-1.0..=1.0).contains(&v),
"Feature '{}' at index {} has value {} outside [-1, 1]",
name,
i,
v,
);
}
}
}
#[test]
fn binary_month_start_only_day_1() {
let ts: Vec<DateTime<Utc>> = (0..91)
.map(|i| Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap() + Duration::days(i as i64))
.collect();
let gen = FeatureGenerator::new().binary(BinaryIndicator::MonthStart);
let feats = gen.generate(&ts);
let col = &feats["month_start"];
let ones: Vec<usize> = col
.iter()
.enumerate()
.filter(|(_, &v)| v == 1.0)
.map(|(i, _)| i)
.collect();
assert_eq!(
ones,
vec![0, 31, 60],
"MonthStart should only fire on day 1"
);
let zero_count = col.iter().filter(|&&v| v == 0.0).count();
assert_eq!(zero_count, 91 - 3);
}
#[test]
fn binary_month_end_various_months() {
let dates = vec![
Utc.with_ymd_and_hms(2024, 1, 31, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 1, 30, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 2, 29, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 2, 28, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2023, 2, 28, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 4, 30, 0, 0, 0).unwrap(),
Utc.with_ymd_and_hms(2024, 4, 29, 0, 0, 0).unwrap(),
];
let gen = FeatureGenerator::new().binary(BinaryIndicator::MonthEnd);
let feats = gen.generate(&dates);
let col = &feats["month_end"];
assert_eq!(col[0], 1.0, "Jan 31 should be month end");
assert_eq!(col[1], 0.0, "Jan 30 should NOT be month end");
assert_eq!(col[2], 1.0, "Feb 29 (leap) should be month end");
assert_eq!(col[3], 0.0, "Feb 28 (leap year) should NOT be month end");
assert_eq!(col[4], 1.0, "Feb 28 (non-leap) should be month end");
assert_eq!(col[5], 1.0, "Apr 30 should be month end");
assert_eq!(col[6], 0.0, "Apr 29 should NOT be month end");
}
#[test]
fn binary_quarter_start() {
let dates = vec![
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 1, 2, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 5, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 7, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 10, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 10, 2, 0, 0, 0).unwrap(), ];
let gen = FeatureGenerator::new().binary(BinaryIndicator::QuarterStart);
let feats = gen.generate(&dates);
let col = &feats["quarter_start"];
assert_eq!(col, &[1.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0]);
}
#[test]
fn binary_quarter_end() {
let dates = vec![
Utc.with_ymd_and_hms(2024, 3, 31, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 3, 30, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 6, 30, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 6, 29, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 9, 30, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 12, 31, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 12, 30, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 1, 31, 0, 0, 0).unwrap(), ];
let gen = FeatureGenerator::new().binary(BinaryIndicator::QuarterEnd);
let feats = gen.generate(&dates);
let col = &feats["quarter_end"];
assert_eq!(col, &[1.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0, 0.0]);
}
#[test]
fn binary_year_start() {
let dates = vec![
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 1, 2, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 12, 31, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2025, 1, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 7, 1, 0, 0, 0).unwrap(), ];
let gen = FeatureGenerator::new().binary(BinaryIndicator::YearStart);
let feats = gen.generate(&dates);
let col = &feats["year_start"];
assert_eq!(col, &[1.0, 0.0, 0.0, 1.0, 0.0]);
}
#[test]
fn binary_year_end() {
let dates = vec![
Utc.with_ymd_and_hms(2024, 12, 31, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 12, 30, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2023, 12, 31, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 6, 30, 0, 0, 0).unwrap(), ];
let gen = FeatureGenerator::new().binary(BinaryIndicator::YearEnd);
let feats = gen.generate(&dates);
let col = &feats["year_end"];
assert_eq!(col, &[1.0, 0.0, 0.0, 1.0, 0.0]);
}
#[test]
fn binary_weekend() {
let ts = daily_timestamps(7);
let gen = FeatureGenerator::new().binary(BinaryIndicator::Weekend);
let feats = gen.generate(&ts);
let col = &feats["weekend"];
assert_eq!(col, &[0.0, 0.0, 0.0, 0.0, 0.0, 1.0, 1.0]);
}
#[test]
fn advanced_days_in_month() {
let dates = vec![
Utc.with_ymd_and_hms(2024, 2, 15, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2023, 2, 15, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 1, 10, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2024, 4, 20, 0, 0, 0).unwrap(), ];
let gen = FeatureGenerator::new().advanced(AdvancedFeature::DaysInMonth);
let feats = gen.generate(&dates);
let col = &feats["days_in_month"];
assert_eq!(col[0], 29.0, "Feb 2024 (leap) should have 29 days");
assert_eq!(col[1], 28.0, "Feb 2023 (non-leap) should have 28 days");
assert_eq!(col[2], 31.0, "Jan should have 31 days");
assert_eq!(col[3], 30.0, "Apr should have 30 days");
}
#[test]
fn advanced_leap_year() {
let dates = vec![
Utc.with_ymd_and_hms(2024, 6, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2023, 6, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(2000, 6, 1, 0, 0, 0).unwrap(), Utc.with_ymd_and_hms(1900, 6, 1, 0, 0, 0).unwrap(), ];
let gen = FeatureGenerator::new().advanced(AdvancedFeature::LeapYear);
let feats = gen.generate(&dates);
let col = &feats["leap_year"];
assert_eq!(col[0], 1.0, "2024 should be leap year");
assert_eq!(col[1], 0.0, "2023 should NOT be leap year");
assert_eq!(col[2], 1.0, "2000 should be leap year");
assert_eq!(col[3], 0.0, "1900 should NOT be leap year");
}
#[test]
fn empty_timestamp_slice() {
let empty: Vec<DateTime<Utc>> = vec![];
let gen = FeatureGenerator::new()
.cyclical(TimeComponent::Month)
.cyclical(TimeComponent::Hour)
.binary(BinaryIndicator::MonthStart)
.binary(BinaryIndicator::Weekend)
.advanced(AdvancedFeature::DaysInMonth)
.advanced(AdvancedFeature::LeapYear)
.fourier(7, 2)
.day_of_week()
.month_of_year();
let feats = gen.generate(&empty);
assert!(
!feats.is_empty(),
"Features map should have keys even for empty input"
);
for (name, values) in &feats {
assert!(
values.is_empty(),
"Feature '{}' should have 0 elements for empty timestamps, got {}",
name,
values.len(),
);
}
}
#[test]
fn single_timestamp() {
let single = vec![Utc.with_ymd_and_hms(2024, 6, 15, 12, 30, 45).unwrap()];
let gen = FeatureGenerator::new()
.cyclical(TimeComponent::Month)
.cyclical(TimeComponent::DayOfWeek)
.cyclical(TimeComponent::Hour)
.binary(BinaryIndicator::MonthStart)
.binary(BinaryIndicator::MonthEnd)
.binary(BinaryIndicator::Weekend)
.advanced(AdvancedFeature::DaysInMonth)
.advanced(AdvancedFeature::LeapYear);
let feats = gen.generate(&single);
for (name, values) in &feats {
assert_eq!(
values.len(),
1,
"Feature '{}' should have 1 element, got {}",
name,
values.len(),
);
}
assert_eq!(feats["weekend"][0], 1.0, "June 15, 2024 is a Saturday");
assert_eq!(feats["month_start"][0], 0.0, "Day 15 is not month start");
assert_eq!(feats["month_end"][0], 0.0, "Day 15 is not month end");
assert_eq!(feats["days_in_month"][0], 30.0, "June has 30 days");
assert_eq!(feats["leap_year"][0], 1.0, "2024 is a leap year");
let expected_sin = (2.0 * PI * 6.0 / 12.0).sin();
let expected_cos = (2.0 * PI * 6.0 / 12.0).cos();
assert!(
(feats["month_sin"][0] - expected_sin).abs() < 1e-10,
"Month sin for June: expected {}, got {}",
expected_sin,
feats["month_sin"][0],
);
assert!(
(feats["month_cos"][0] - expected_cos).abs() < 1e-10,
"Month cos for June: expected {}, got {}",
expected_cos,
feats["month_cos"][0],
);
}
}