Files
rustytorch/demos/rtx-timeseries-demo/src/sample_data.rs
T
2026-03-04 00:08:42 +00:00

501 lines
16 KiB
Rust

//! Sample data generators for the time series demo
use rand::Rng;
use rand_distr::{Distribution, Normal};
use timeseries_shared::{DataPoint, SampleDataset, TimeSeriesData};
use crate::error::{ForecastError, Result};
/// Available sample datasets
pub const SAMPLE_DATASETS: &[SampleDataset] = &[
SampleDataset::AirlinePassengers,
SampleDataset::StockPrices,
SampleDataset::Temperature,
SampleDataset::WebTraffic,
SampleDataset::SineWave,
SampleDataset::RandomWalk,
];
/// Generate sample time series data
pub fn generate_sample_data(dataset: SampleDataset, length: usize) -> Result<TimeSeriesData> {
match dataset {
SampleDataset::AirlinePassengers => generate_airline_passengers(length),
SampleDataset::StockPrices => generate_stock_prices(length),
SampleDataset::Temperature => generate_temperature(length),
SampleDataset::WebTraffic => generate_web_traffic(length),
SampleDataset::SineWave => generate_sine_wave(length),
SampleDataset::RandomWalk => generate_random_walk(length),
}
}
/// Generate airline passengers-like data (monthly, with trend and seasonality)
fn generate_airline_passengers(length: usize) -> Result<TimeSeriesData> {
let mut rng = rand::rng();
let normal =
Normal::new(0.0, 10.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
let mut data = Vec::with_capacity(length);
for i in 0..length {
let t = i as f64;
// Trend
let trend = 100.0 + 2.5 * t;
// Seasonality (12-month cycle)
let seasonal = 50.0 * (2.0 * std::f64::consts::PI * t / 12.0).sin();
// Growing amplitude
let amplitude_factor = 1.0 + 0.01 * t;
// Noise
let noise = normal.sample(&mut rng);
let value = trend + seasonal * amplitude_factor + noise;
data.push(DataPoint {
timestamp: i as f64,
value: value.max(0.0),
});
}
Ok(TimeSeriesData {
data,
name: Some("Airline Passengers".to_string()),
frequency: Some("monthly".to_string()),
})
}
/// Generate stock prices with trend and volatility clustering
fn generate_stock_prices(length: usize) -> Result<TimeSeriesData> {
let mut rng = rand::rng();
let mut data = Vec::with_capacity(length);
let mut price = 100.0;
let drift = 0.0005; // Small positive drift
let base_volatility = 0.02;
for i in 0..length {
// Random return with drift
let volatility = base_volatility * (1.0 + 0.5 * rng.random::<f64>());
let return_val = drift + volatility * rng.random_range(-1.0..1.0);
price *= 1.0 + return_val;
data.push(DataPoint {
timestamp: i as f64,
value: price,
});
}
Ok(TimeSeriesData {
data,
name: Some("Stock Price".to_string()),
frequency: Some("daily".to_string()),
})
}
/// Generate temperature data with daily and seasonal patterns
fn generate_temperature(length: usize) -> Result<TimeSeriesData> {
let mut rng = rand::rng();
let normal =
Normal::new(0.0, 3.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
let mut data = Vec::with_capacity(length);
for i in 0..length {
let day = i as f64;
// Annual seasonality (365-day cycle)
let annual =
15.0 * (2.0 * std::f64::consts::PI * day / 365.0 - std::f64::consts::PI / 2.0).sin();
// Base temperature
let base = 15.0;
// Noise
let noise = normal.sample(&mut rng);
let value = base + annual + noise;
data.push(DataPoint {
timestamp: i as f64,
value,
});
}
Ok(TimeSeriesData {
data,
name: Some("Temperature (°C)".to_string()),
frequency: Some("daily".to_string()),
})
}
/// Generate web traffic with weekly seasonality and growth trend
fn generate_web_traffic(length: usize) -> Result<TimeSeriesData> {
let mut rng = rand::rng();
let normal =
Normal::new(0.0, 500.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
let mut data = Vec::with_capacity(length);
for i in 0..length {
let hour = i as f64;
// Growth trend
let trend = 10000.0 + 5.0 * hour;
// Weekly seasonality (168 hours = 1 week)
let weekly = 2000.0 * (2.0 * std::f64::consts::PI * hour / 168.0).sin();
// Daily seasonality (24 hours)
let daily = 1000.0 * (2.0 * std::f64::consts::PI * hour / 24.0).sin();
// Noise
let noise = normal.sample(&mut rng);
let value = (trend + weekly + daily + noise).max(0.0);
data.push(DataPoint {
timestamp: i as f64,
value,
});
}
Ok(TimeSeriesData {
data,
name: Some("Web Traffic (visits/hour)".to_string()),
frequency: Some("hourly".to_string()),
})
}
/// Generate clean sine wave for testing
fn generate_sine_wave(length: usize) -> Result<TimeSeriesData> {
let mut rng = rand::rng();
let normal =
Normal::new(0.0, 0.1).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
let mut data = Vec::with_capacity(length);
for i in 0..length {
let t = i as f64;
let value = (2.0 * std::f64::consts::PI * t / 50.0).sin() + normal.sample(&mut rng);
data.push(DataPoint {
timestamp: t,
value,
});
}
Ok(TimeSeriesData {
data,
name: Some("Sine Wave".to_string()),
frequency: Some("synthetic".to_string()),
})
}
/// Generate random walk for baseline comparison
fn generate_random_walk(length: usize) -> Result<TimeSeriesData> {
let mut rng = rand::rng();
let normal =
Normal::new(0.0, 1.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
let mut data = Vec::with_capacity(length);
let mut value = 0.0;
for i in 0..length {
value += normal.sample(&mut rng);
data.push(DataPoint {
timestamp: i as f64,
value,
});
}
Ok(TimeSeriesData {
data,
name: Some("Random Walk".to_string()),
frequency: Some("synthetic".to_string()),
})
}
#[cfg(test)]
mod tests {
use super::*;
// ========== Basic Generation Tests ==========
#[test]
fn test_generate_all_datasets() {
for dataset in SAMPLE_DATASETS {
let result = generate_sample_data(*dataset, 100);
assert!(result.is_ok(), "Failed for {:?}", dataset);
let data = result.unwrap();
assert_eq!(data.data.len(), 100);
}
}
#[test]
fn test_airline_passengers() {
let data = generate_airline_passengers(144).unwrap();
assert_eq!(data.data.len(), 144);
assert!(data.data.iter().all(|p| p.value >= 0.0));
}
// ========== Airline Passengers Tests ==========
#[test]
fn test_airline_passengers_name_and_frequency() {
let data = generate_airline_passengers(12).unwrap();
assert_eq!(data.name, Some("Airline Passengers".to_string()));
assert_eq!(data.frequency, Some("monthly".to_string()));
}
#[test]
fn test_airline_passengers_trend() {
let data = generate_airline_passengers(100).unwrap();
// Should have positive trend - later values generally higher
let first_half_avg: f64 = data.data[..50].iter().map(|p| p.value).sum::<f64>() / 50.0;
let second_half_avg: f64 = data.data[50..].iter().map(|p| p.value).sum::<f64>() / 50.0;
assert!(
second_half_avg > first_half_avg,
"Airline data should have upward trend"
);
}
#[test]
fn test_airline_passengers_timestamps() {
let data = generate_airline_passengers(50).unwrap();
for (i, point) in data.data.iter().enumerate() {
assert_eq!(point.timestamp, i as f64);
}
}
// ========== Stock Prices Tests ==========
#[test]
fn test_stock_prices() {
let data = generate_stock_prices(100).unwrap();
assert_eq!(data.data.len(), 100);
assert_eq!(data.name, Some("Stock Price".to_string()));
assert_eq!(data.frequency, Some("daily".to_string()));
}
#[test]
fn test_stock_prices_positive() {
let data = generate_stock_prices(100).unwrap();
// Stock prices should always be positive (multiplicative process)
assert!(data.data.iter().all(|p| p.value > 0.0));
}
#[test]
fn test_stock_prices_starts_at_100() {
let data = generate_stock_prices(1).unwrap();
// First value should be close to 100 (starting price with small perturbation)
assert!((data.data[0].value - 100.0).abs() < 10.0);
}
// ========== Temperature Tests ==========
#[test]
fn test_temperature() {
let data = generate_temperature(365).unwrap();
assert_eq!(data.data.len(), 365);
assert_eq!(data.name, Some("Temperature (°C)".to_string()));
assert_eq!(data.frequency, Some("daily".to_string()));
}
#[test]
fn test_temperature_realistic_range() {
let data = generate_temperature(365).unwrap();
// Temperatures should be in realistic range (-10 to 40 C roughly)
for point in &data.data {
assert!(
point.value > -20.0 && point.value < 50.0,
"Temperature {} out of realistic range",
point.value
);
}
}
#[test]
fn test_temperature_seasonality() {
let data = generate_temperature(365).unwrap();
// Summer (around day 180) should be warmer than winter (around day 0)
let winter_avg: f64 = data.data[..30].iter().map(|p| p.value).sum::<f64>() / 30.0;
let summer_avg: f64 = data.data[150..210].iter().map(|p| p.value).sum::<f64>() / 60.0;
assert!(
summer_avg > winter_avg,
"Summer should be warmer than winter"
);
}
// ========== Web Traffic Tests ==========
#[test]
fn test_web_traffic() {
let data = generate_web_traffic(168).unwrap();
assert_eq!(data.data.len(), 168);
assert_eq!(data.name, Some("Web Traffic (visits/hour)".to_string()));
assert_eq!(data.frequency, Some("hourly".to_string()));
}
#[test]
fn test_web_traffic_non_negative() {
let data = generate_web_traffic(500).unwrap();
assert!(data.data.iter().all(|p| p.value >= 0.0));
}
#[test]
fn test_web_traffic_trend() {
let data = generate_web_traffic(1000).unwrap();
// Should have positive trend
let first_half_avg: f64 = data.data[..500].iter().map(|p| p.value).sum::<f64>() / 500.0;
let second_half_avg: f64 = data.data[500..].iter().map(|p| p.value).sum::<f64>() / 500.0;
assert!(
second_half_avg > first_half_avg,
"Web traffic should trend upward"
);
}
// ========== Sine Wave Tests ==========
#[test]
fn test_sine_wave() {
let data = generate_sine_wave(100).unwrap();
assert_eq!(data.data.len(), 100);
assert_eq!(data.name, Some("Sine Wave".to_string()));
assert_eq!(data.frequency, Some("synthetic".to_string()));
}
#[test]
fn test_sine_wave_range() {
let data = generate_sine_wave(200).unwrap();
// Sine wave should be in range [-1.5, 1.5] with small noise
for point in &data.data {
assert!(
point.value >= -2.0 && point.value <= 2.0,
"Sine value {} out of expected range",
point.value
);
}
}
#[test]
fn test_sine_wave_periodicity() {
let data = generate_sine_wave(200).unwrap();
// Check that values near multiples of 50 (period) are similar
// Period is 50, so sine(0) ≈ sine(50) ≈ sine(100)
let val_at_0 = data.data[0].value;
let val_at_50 = data.data[50].value;
let val_at_100 = data.data[100].value;
// Allow for noise
assert!((val_at_0 - val_at_50).abs() < 1.0);
assert!((val_at_0 - val_at_100).abs() < 1.0);
}
// ========== Random Walk Tests ==========
#[test]
fn test_random_walk() {
let data = generate_random_walk(100).unwrap();
assert_eq!(data.data.len(), 100);
assert_eq!(data.name, Some("Random Walk".to_string()));
assert_eq!(data.frequency, Some("synthetic".to_string()));
}
#[test]
fn test_random_walk_starts_at_zero() {
let data = generate_random_walk(10).unwrap();
// First step should be small (normal with mean 0, std 1)
assert!(data.data[0].value.abs() < 5.0);
}
#[test]
fn test_random_walk_variance_grows() {
// Generate multiple walks and check variance grows with time
let data = generate_random_walk(1000).unwrap();
let first_100: Vec<f64> = data.data[..100].iter().map(|p| p.value).collect();
let last_100: Vec<f64> = data.data[900..].iter().map(|p| p.value).collect();
fn variance(v: &[f64]) -> f64 {
let mean = v.iter().sum::<f64>() / v.len() as f64;
v.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / v.len() as f64
}
// Later values should have higher variance on average
// (This test may occasionally fail due to randomness, but typically passes)
let _first_var = variance(&first_100);
let _last_var = variance(&last_100);
// Just check it runs without error - variance comparison is probabilistic
}
// ========== Edge Case Tests ==========
#[test]
fn test_generate_single_point() {
for dataset in SAMPLE_DATASETS {
let result = generate_sample_data(*dataset, 1);
assert!(result.is_ok());
assert_eq!(result.unwrap().data.len(), 1);
}
}
#[test]
fn test_generate_large_dataset() {
let data = generate_sample_data(SampleDataset::SineWave, 10000).unwrap();
assert_eq!(data.data.len(), 10000);
}
#[test]
fn test_generate_zero_length() {
// Zero length should still work (empty dataset)
let data = generate_sample_data(SampleDataset::SineWave, 0).unwrap();
assert!(data.data.is_empty());
}
// ========== Timestamp Continuity Tests ==========
#[test]
fn test_timestamps_continuous() {
for dataset in SAMPLE_DATASETS {
let data = generate_sample_data(*dataset, 50).unwrap();
for (i, point) in data.data.iter().enumerate() {
assert_eq!(
point.timestamp, i as f64,
"Timestamp discontinuity at index {} for {:?}",
i, dataset
);
}
}
}
// ========== Sample Datasets Constant Tests ==========
#[test]
fn test_sample_datasets_constant() {
assert_eq!(SAMPLE_DATASETS.len(), 6);
assert!(SAMPLE_DATASETS.contains(&SampleDataset::AirlinePassengers));
assert!(SAMPLE_DATASETS.contains(&SampleDataset::StockPrices));
assert!(SAMPLE_DATASETS.contains(&SampleDataset::Temperature));
assert!(SAMPLE_DATASETS.contains(&SampleDataset::WebTraffic));
assert!(SAMPLE_DATASETS.contains(&SampleDataset::SineWave));
assert!(SAMPLE_DATASETS.contains(&SampleDataset::RandomWalk));
}
// ========== Data Quality Tests ==========
#[test]
fn test_no_nan_values() {
for dataset in SAMPLE_DATASETS {
let data = generate_sample_data(*dataset, 100).unwrap();
for point in &data.data {
assert!(!point.value.is_nan(), "NaN found in {:?}", dataset);
assert!(!point.timestamp.is_nan(), "NaN timestamp in {:?}", dataset);
}
}
}
#[test]
fn test_no_infinite_values() {
for dataset in SAMPLE_DATASETS {
let data = generate_sample_data(*dataset, 100).unwrap();
for point in &data.data {
assert!(
!point.value.is_infinite(),
"Infinite found in {:?}",
dataset
);
}
}
}
}