Files
rustytorch/crates/models/rtx-timeseries/tests/arima_tests.rs
T
2026-03-04 00:08:42 +00:00

483 lines
16 KiB
Rust

//! Comprehensive tests for ARIMA models
use approx::assert_abs_diff_eq;
use rtx_tensor::{Device, Tensor};
use rtx_timeseries::{
TimeSeriesError,
models::{ARIMAConfig, ARIMAModel, TimeSeriesModel, TypedTimeSeriesModel},
};
use tokio_test;
#[tokio::test]
async fn test_arima_model_creation() -> Result<(), Box<dyn std::error::Error>> {
let model = ARIMAModel::new((1, 1, 1), None);
assert_eq!(model.get_config().order, (1, 1, 1));
assert!(model.is_fitted().is_err());
Ok(())
}
#[tokio::test]
async fn test_arima_model_with_custom_config() -> Result<(), Box<dyn std::error::Error>> {
let config = ARIMAConfig {
order: (2, 1, 2),
include_constant: true,
max_iter: 500,
tolerance: 1e-6,
quantum_enhanced: false,
device: "cpu".to_string(),
};
let model = ARIMAModel::with_config(config.clone());
assert_eq!(model.get_config().order, (2, 1, 2));
assert_eq!(model.get_config().max_iter, 500);
assert!(!model.get_config().quantum_enhanced);
Ok(())
}
#[tokio::test]
async fn test_arima_parameters_validation() -> Result<(), Box<dyn std::error::Error>> {
use rtx_timeseries::models::ARIMAParameters;
let params = ARIMAParameters::new(vec![0.5, -0.2], vec![0.3], 0.1, 1.0);
assert_eq!(params.param_count(), 5); // 2 AR + 1 MA + const + sigma2
assert!(params.is_stationary());
assert!(params.is_invertible());
Ok(())
}
#[tokio::test]
async fn test_arima_parameters_non_stationary() -> Result<(), Box<dyn std::error::Error>> {
use rtx_timeseries::models::ARIMAParameters;
let params = ARIMAParameters::new(
vec![0.8, 0.3], // Sum > 1, not stationary
vec![0.2],
0.0,
1.0,
);
assert!(!params.is_stationary());
assert!(params.is_invertible());
Ok(())
}
#[tokio::test]
async fn test_arima_fitting_simple_trend() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Create simple linear trend data
let data_vec: Vec<f32> = (0..50).map(|i| i as f32 * 0.5 + 10.0).collect();
let data = Tensor::from_vec(data_vec, &[50], &device)?;
let timestamps = Tensor::arange(0, 50, &device)?;
let mut model = ARIMAModel::new((1, 1, 1), None);
let result = TimeSeriesModel::fit(&mut model, &data, &timestamps).await;
assert!(result.is_ok(), "ARIMA fitting failed: {:?}", result.err());
assert!(model.is_fitted().is_ok());
// Check that parameters were estimated
let params = model.get_parameters().unwrap();
assert!(params.contains_key("ar.L1"));
assert!(params.contains_key("ma.L1"));
assert!(params.contains_key("const"));
assert!(params.contains_key("sigma2"));
Ok(())
}
#[tokio::test]
async fn test_arima_fitting_seasonal_data() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Create data with trend and seasonality
let mut data_vec = Vec::new();
for i in 0..100 {
let t = i as f32;
let trend = 0.1 * t;
let seasonal = 2.0 * (2.0 * std::f32::consts::PI * t / 12.0).sin();
let noise = 0.1 * (rand::random::<f32>() - 0.5);
data_vec.push(trend + seasonal + noise + 50.0);
}
let data = Tensor::from_vec(data_vec, &[100], &device)?;
let timestamps = Tensor::arange(0, 100, &device)?;
let mut model = ARIMAModel::new((2, 1, 2), None);
let result = TimeSeriesModel::fit(&mut model, &data, &timestamps).await;
assert!(result.is_ok());
assert!(model.is_fitted().is_ok());
// Verify fit quality
let fit_metrics = model
.calculate_fit_metrics(&data, &timestamps)
.await
.unwrap();
assert!(fit_metrics.r_squared > 0.0); // Should capture some variance
assert!(fit_metrics.rmse < 10.0); // Reasonable error for this data
Ok(())
}
#[tokio::test]
async fn test_arima_forecasting() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Create predictable data
let data_vec: Vec<f32> = (1..=20).map(|i| i as f32).collect();
let data = Tensor::from_vec(data_vec, &[20], &device)?;
let timestamps = Tensor::arange(1, 21, &device)?;
let mut model = ARIMAModel::new((1, 0, 0), None);
TimeSeriesModel::fit(&mut model, &data, &timestamps)
.await
.unwrap();
// Generate forecasts
let forecast = model.forecast(5, 0.95).await.unwrap();
assert_eq!(forecast.len(), 5);
assert_eq!(forecast.mean.shape()[0], 5);
assert_eq!(forecast.lower.shape()[0], 5);
assert_eq!(forecast.upper.shape()[0], 5);
// Check that confidence intervals are reasonable
for i in 0..5 {
let mean_val = forecast.mean.get(&[i]).unwrap();
let lower_val = forecast.lower.get(&[i]).unwrap();
let upper_val = forecast.upper.get(&[i]).unwrap();
assert!(lower_val < mean_val);
assert!(mean_val < upper_val);
assert!(upper_val - lower_val > 0.0); // Non-zero uncertainty
}
Ok(())
}
#[tokio::test]
async fn test_arima_forecasting_different_confidence_levels()
-> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
let data = Tensor::from_vec(
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0],
&[10],
&device,
)?;
let timestamps = Tensor::arange(1, 11, &device)?;
let mut model = ARIMAModel::new((1, 0, 0), None);
TimeSeriesModel::fit(&mut model, &data, &timestamps)
.await
.unwrap();
// Test different confidence levels
for &confidence_level in &[0.90, 0.95, 0.99] {
let forecast = model.forecast(3, confidence_level).await.unwrap();
for i in 0..3 {
let lower_val = forecast.lower.get(&[i]).unwrap();
let upper_val = forecast.upper.get(&[i]).unwrap();
let interval_width = upper_val - lower_val;
// Higher confidence should give wider intervals
assert!(interval_width > 0.0);
if confidence_level == 0.99 {
// 99% intervals should be wider than others
let forecast_95 = model.forecast(3, 0.95).await.unwrap();
let width_95 =
forecast_95.upper.get(&[i]).unwrap() - forecast_95.lower.get(&[i]).unwrap();
assert!(interval_width >= width_95);
}
}
}
Ok(())
}
#[tokio::test]
async fn test_arima_residuals() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
let data = Tensor::from_vec(
vec![1.0, 2.1, 2.9, 4.1, 4.9, 6.1, 6.9, 8.1, 8.9, 10.1],
&[10],
&device,
)?;
let timestamps = Tensor::arange(0, 10, &device)?;
let mut model = ARIMAModel::new((1, 0, 0), None);
TimeSeriesModel::fit(&mut model, &data, &timestamps)
.await
.unwrap();
let residuals = model.residuals(&data, &timestamps).await.unwrap();
assert_eq!(residuals.shape()[0], 10);
// Check that residuals are reasonable (should be small for this nearly linear data)
let residuals_vec: Vec<f32> = (0..10).map(|i| residuals.get(&[i]).unwrap()).collect();
let mean_abs_residual = residuals_vec.iter().map(|r| r.abs()).sum::<f32>() / 10.0;
assert!(mean_abs_residual < 1.0); // Should be small for this data
Ok(())
}
#[tokio::test]
async fn test_arima_parameter_getters_setters() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
let data = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0], &[5], &device)?;
let timestamps = Tensor::arange(0, 5, &device)?;
let mut model = ARIMAModel::new((1, 0, 1), None);
TimeSeriesModel::fit(&mut model, &data, &timestamps)
.await
.unwrap();
// Test parameter getting
let params = model.get_parameters().unwrap();
assert!(params.contains_key("ar.L1"));
assert!(params.contains_key("ma.L1"));
assert!(params.contains_key("const"));
assert!(params.contains_key("sigma2"));
// Test parameter setting
let mut new_params = std::collections::HashMap::new();
new_params.insert("ar.L1".to_string(), 0.5);
new_params.insert("ma.L1".to_string(), 0.3);
new_params.insert("const".to_string(), 1.0);
new_params.insert("sigma2".to_string(), 0.8);
let result = model.set_parameters(new_params);
assert!(result.is_ok());
// Verify parameters were set
let updated_params = model.get_parameters().unwrap();
assert_abs_diff_eq!(updated_params["ar.L1"], 0.5, epsilon = 1e-6);
assert_abs_diff_eq!(updated_params["ma.L1"], 0.3, epsilon = 1e-6);
assert_abs_diff_eq!(updated_params["const"], 1.0, epsilon = 1e-6);
assert_abs_diff_eq!(updated_params["sigma2"], 0.8, epsilon = 1e-6);
Ok(())
}
#[tokio::test]
async fn test_arima_fit_metrics() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Create data with known properties
let data_vec: Vec<f32> = (0..50)
.map(|i| i as f32 * 0.5 + 10.0 + 0.1 * (rand::random::<f32>() - 0.5))
.collect();
let data = Tensor::from_vec(data_vec, &[50], &device)?;
let timestamps = Tensor::arange(0, 50, &device)?;
let mut model = ARIMAModel::new((1, 1, 1), None);
TimeSeriesModel::fit(&mut model, &data, &timestamps)
.await
.unwrap();
let metrics = model
.calculate_fit_metrics(&data, &timestamps)
.await
.unwrap();
// Check that metrics are reasonable
assert!(metrics.aic.is_finite());
assert!(metrics.bic.is_finite());
assert!(metrics.r_squared >= 0.0 && metrics.r_squared <= 1.0);
assert!(metrics.rmse >= 0.0);
assert!(metrics.mae >= 0.0);
assert!(metrics.mape >= 0.0);
// For trend data, should have decent R²
assert!(metrics.r_squared > 0.5);
Ok(())
}
#[tokio::test]
async fn test_arima_model_serialization() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
let data = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0], &[5], &device)?;
let timestamps = Tensor::arange(0, 5, &device)?;
let mut model = ARIMAModel::new((1, 0, 1), None);
TimeSeriesModel::fit(&mut model, &data, &timestamps)
.await
.unwrap();
// Test serialization
let serialized = model.serialize();
assert!(serialized.is_ok());
// Test deserialization
let mut new_model = ARIMAModel::new((1, 0, 1), None);
let deserialization_result = new_model.deserialize(&serialized.unwrap());
assert!(deserialization_result.is_ok());
// Verify models are equivalent
assert!(new_model.is_fitted().is_ok());
let original_params = model.get_parameters().unwrap();
let deserialized_params = new_model.get_parameters().unwrap();
for (key, value) in original_params {
assert_abs_diff_eq!(deserialized_params[&key], value, epsilon = 1e-6);
}
Ok(())
}
#[tokio::test]
async fn test_arima_model_cloning() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
let data = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0], &[5], &device)?;
let timestamps = Tensor::arange(0, 5, &device)?;
let mut model = ARIMAModel::new((1, 0, 1), None);
TimeSeriesModel::fit(&mut model, &data, &timestamps)
.await
.unwrap();
// Test model cloning
let cloned_model = model.clone_model();
assert!(cloned_model.is_ok());
let cloned = cloned_model.unwrap();
assert!(cloned.is_fitted().is_ok());
// Verify parameters match
let original_params = model.get_parameters().unwrap();
let cloned_params = cloned.get_parameters().unwrap();
assert_eq!(original_params.len(), cloned_params.len());
for (key, value) in original_params {
assert_abs_diff_eq!(cloned_params[&key], value, epsilon = 1e-6);
}
Ok(())
}
#[tokio::test]
async fn test_arima_validation_errors() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Test mismatched data and timestamps
let data = Tensor::from_vec(vec![1.0, 2.0, 3.0], &[3], &device)?;
let timestamps = Tensor::from_vec(vec![1.0, 2.0], &[2], &device)?; // Wrong length
let mut model = ARIMAModel::new((1, 0, 1), None);
let result = TimeSeriesModel::fit(&mut model, &data, &timestamps).await;
assert!(result.is_err());
match result.err().unwrap() {
TimeSeriesError::ValidationError(_) => {} // Expected
other => panic!("Expected ValidationError, got {:?}", other),
}
Ok(())
}
#[tokio::test]
async fn test_arima_insufficient_data() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Test with insufficient data points
let data = Tensor::from_vec(vec![1.0, 2.0], &[2], &device)?; // Only 2 points
let timestamps = Tensor::from_vec(vec![1.0, 2.0], &[2], &device)?;
let mut model = ARIMAModel::new((2, 1, 2), None); // Requires more data
let result = TimeSeriesModel::fit(&mut model, &data, &timestamps).await;
assert!(result.is_err());
match result.err().unwrap() {
TimeSeriesError::ValidationError(_) => {} // Expected
other => panic!("Expected ValidationError, got {:?}", other),
}
Ok(())
}
#[tokio::test]
async fn test_arima_forecasting_unfitted_model() -> Result<(), Box<dyn std::error::Error>> {
let model = ARIMAModel::new((1, 0, 1), None);
let result = model.forecast(5, 0.95).await;
assert!(result.is_err());
match result.err().unwrap() {
TimeSeriesError::ModelStateError(_) => {} // Expected
other => panic!("Expected ModelStateError, got {:?}", other),
}
Ok(())
}
#[tokio::test]
async fn test_arima_differencing() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Test differencing operation
let data = Tensor::from_vec(vec![1.0, 3.0, 6.0, 10.0, 15.0], &[5], &device)?;
let model = ARIMAModel::new((1, 1, 1), None);
// First difference should be [2, 3, 4, 5]
let diff = model.difference_series(&data, 1).await.unwrap();
assert_eq!(diff.shape()[0], 4);
let expected_diff = vec![2.0, 3.0, 4.0, 5.0];
for (i, expected) in expected_diff.iter().enumerate() {
let actual = diff.get(&[i]).unwrap();
assert_abs_diff_eq!(actual, *expected, epsilon = 1e-6);
}
// Second difference should be [1, 1, 1]
let diff2 = model.difference_series(&data, 2).await.unwrap();
assert_eq!(diff2.shape()[0], 3);
for i in 0..3 {
let actual = diff2.get(&[i]).unwrap();
assert_abs_diff_eq!(actual, 1.0, epsilon = 1e-6);
}
Ok(())
}
#[tokio::test]
async fn test_arima_integration() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
// Test integration (reverse differencing)
let original = Tensor::from_vec(vec![1.0, 3.0, 6.0, 10.0, 15.0], &[5], &device)?;
let model = ARIMAModel::new((1, 1, 1), None);
// Difference and then integrate should recover original (approximately)
let differenced = model.difference_series(&original, 1).await.unwrap();
let integrated = model
.integrate_series(&differenced, &original, 1)
.await
.unwrap();
// Should recover original series (may have one extra point)
let min_len = original.shape()[0].min(integrated.shape()[0]);
for i in 0..min_len {
let original_val = original.get(&[i]).unwrap();
let integrated_val = integrated.get(&[i]).unwrap();
assert_abs_diff_eq!(integrated_val, original_val, epsilon = 1e-6);
}
Ok(())
}
#[tokio::test]
async fn test_arima_zero_differencing() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cpu();
let data = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0], &[5], &device)?;
let model = ARIMAModel::new((1, 0, 1), None);
// Zero differencing should return original data
let result = model.difference_series(&data, 0).await.unwrap();
assert_eq!(result.shape()[0], data.shape()[0]);
for i in 0..5 {
let original_val = data.get(&[i]).unwrap();
let result_val = result.get(&[i]).unwrap();
assert_abs_diff_eq!(result_val, original_val, epsilon = 1e-6);
}
Ok(())
}