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

251 lines
7.3 KiB
Rust

//! Error types for the time series library
use thiserror::Error;
/// Result type for time series operations
pub type Result<T> = std::result::Result<T, TimeSeriesError>;
/// Comprehensive error types for time series analysis and forecasting
#[derive(Error, Debug)]
pub enum TimeSeriesError {
/// Input validation errors
#[error("Validation error: {0}")]
ValidationError(String),
/// Model fitting errors
#[error("Model fitting error: {0}")]
FittingError(String),
/// Forecasting errors
#[error("Forecasting error: {0}")]
ForecastingError(String),
/// Data preprocessing errors
#[error("Data preprocessing error: {0}")]
PreprocessingError(String),
/// Configuration errors
#[error("Configuration error: {0}")]
ConfigurationError(String),
/// Serialization/deserialization errors
#[error("Serialization error: {0}")]
SerializationError(String),
/// GPU/device errors
#[error("Device error: {0}")]
DeviceError(String),
/// Optimization convergence errors
#[error("Optimization error: {0}")]
OptimizationError(String),
/// Stationarity test errors
#[error("Stationarity test error: {0}")]
StationarityError(String),
/// Seasonality detection errors
#[error("Seasonality detection error: {0}")]
SeasonalityError(String),
/// Anomaly detection errors
#[error("Anomaly detection error: {0}")]
AnomalyDetectionError(String),
/// Memory allocation errors
#[error("Memory error: {0}")]
MemoryError(String),
/// Numerical computation errors
#[error("Numerical error: {0}")]
NumericalError(String),
/// Missing data handling errors
#[error("Missing data error: {0}")]
MissingDataError(String),
/// Model state errors
#[error("Model state error: {0}")]
ModelStateError(String),
/// Time dimension errors
#[error("Time dimension error: {0}")]
TimeDimensionError(String),
/// Frequency detection errors
#[error("Frequency detection error: {0}")]
FrequencyError(String),
/// Integration with other RTX components
#[error("Integration error: {0}")]
IntegrationError(String),
/// Tensor operations errors
#[error("Tensor error: {source}")]
TensorError {
#[from]
source: rtx_tensor::TensorError,
},
/// Runtime errors
#[error("Runtime error: {source}")]
RuntimeError {
#[from]
source: rtx_runtime::RuntimeError,
},
/// Validation library errors
#[error("Validation error: {source}")]
ValidationLibraryError {
#[from]
source: rtx_validation::ValidationError,
},
/// AutoML errors
#[error("AutoML error: {source}")]
AutoMLError {
#[from]
source: rtx_automeasure::AutoMLError,
},
/// Generic I/O errors
#[error("I/O error: {source}")]
IoError {
#[from]
source: std::io::Error,
},
/// JSON serialization errors
#[error("JSON error: {source}")]
JsonError {
#[from]
source: serde_json::Error,
},
/// Generic errors
#[error("Internal error: {0}")]
InternalError(String),
}
impl TimeSeriesError {
/// Create a validation error
pub fn validation<S: Into<String>>(msg: S) -> Self {
Self::ValidationError(msg.into())
}
/// Create a fitting error
pub fn fitting<S: Into<String>>(msg: S) -> Self {
Self::FittingError(msg.into())
}
/// Create a forecasting error
pub fn forecasting<S: Into<String>>(msg: S) -> Self {
Self::ForecastingError(msg.into())
}
/// Create a configuration error
pub fn configuration<S: Into<String>>(msg: S) -> Self {
Self::ConfigurationError(msg.into())
}
/// Create a device error
pub fn device<S: Into<String>>(msg: S) -> Self {
Self::DeviceError(msg.into())
}
/// Create an optimization error
pub fn optimization<S: Into<String>>(msg: S) -> Self {
Self::OptimizationError(msg.into())
}
/// Create a numerical error
pub fn numerical<S: Into<String>>(msg: S) -> Self {
Self::NumericalError(msg.into())
}
/// Create a model state error
pub fn model_state<S: Into<String>>(msg: S) -> Self {
Self::ModelStateError(msg.into())
}
/// Check if error is recoverable
pub fn is_recoverable(&self) -> bool {
match self {
Self::ValidationError(_) => false,
Self::ConfigurationError(_) => false,
Self::DeviceError(_) => true,
Self::OptimizationError(_) => true,
Self::MemoryError(_) => true,
Self::NumericalError(_) => true,
Self::IntegrationError(_) => true,
_ => false,
}
}
/// Get error category
pub fn category(&self) -> &'static str {
match self {
Self::ValidationError(_) => "validation",
Self::FittingError(_) => "fitting",
Self::ForecastingError(_) => "forecasting",
Self::PreprocessingError(_) => "preprocessing",
Self::ConfigurationError(_) => "configuration",
Self::SerializationError(_) => "serialization",
Self::DeviceError(_) => "device",
Self::OptimizationError(_) => "optimization",
Self::StationarityError(_) => "stationarity",
Self::SeasonalityError(_) => "seasonality",
Self::AnomalyDetectionError(_) => "anomaly_detection",
Self::MemoryError(_) => "memory",
Self::NumericalError(_) => "numerical",
Self::MissingDataError(_) => "missing_data",
Self::ModelStateError(_) => "model_state",
Self::TimeDimensionError(_) => "time_dimension",
Self::FrequencyError(_) => "frequency",
Self::IntegrationError(_) => "integration",
Self::TensorError { .. } => "tensor",
Self::RuntimeError { .. } => "runtime",
Self::ValidationLibraryError { .. } => "validation_library",
Self::AutoMLError { .. } => "automl",
Self::IoError { .. } => "io",
Self::JsonError { .. } => "json",
Self::InternalError(_) => "internal",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_error_creation() {
let err = TimeSeriesError::validation("Invalid input");
assert_eq!(err.category(), "validation");
assert!(!err.is_recoverable());
let err = TimeSeriesError::device("GPU out of memory");
assert_eq!(err.category(), "device");
assert!(err.is_recoverable());
}
#[test]
fn test_error_display() {
let err = TimeSeriesError::fitting("Model failed to converge");
let display = format!("{}", err);
assert!(display.contains("Model fitting error"));
assert!(display.contains("Model failed to converge"));
}
#[test]
fn test_error_categorization() {
let validation_err = TimeSeriesError::ValidationError("test".to_string());
let device_err = TimeSeriesError::DeviceError("test".to_string());
let numerical_err = TimeSeriesError::NumericalError("test".to_string());
assert_eq!(validation_err.category(), "validation");
assert_eq!(device_err.category(), "device");
assert_eq!(numerical_err.category(), "numerical");
}
}