From 5bc56eb9c8d7bfb6060f7c4c5a54b8a98507d861 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 23 Feb 2026 12:54:36 +0100 Subject: [PATCH] feat(ml): add data_pipeline module with DatasetSpec and DatasetMode types Co-Authored-By: Claude Opus 4.6 --- ml/src/data_pipeline/cache.rs | 1 + ml/src/data_pipeline/manager.rs | 1 + ml/src/data_pipeline/mod.rs | 264 +++++++++++++++++++++++++++++++ ml/src/data_pipeline/prepared.rs | 1 + ml/src/lib.rs | 1 + 5 files changed, 268 insertions(+) create mode 100644 ml/src/data_pipeline/cache.rs create mode 100644 ml/src/data_pipeline/manager.rs create mode 100644 ml/src/data_pipeline/mod.rs create mode 100644 ml/src/data_pipeline/prepared.rs diff --git a/ml/src/data_pipeline/cache.rs b/ml/src/data_pipeline/cache.rs new file mode 100644 index 000000000..e8b560b03 --- /dev/null +++ b/ml/src/data_pipeline/cache.rs @@ -0,0 +1 @@ +//! Cache module - TBD diff --git a/ml/src/data_pipeline/manager.rs b/ml/src/data_pipeline/manager.rs new file mode 100644 index 000000000..0a5ed1dcd --- /dev/null +++ b/ml/src/data_pipeline/manager.rs @@ -0,0 +1 @@ +//! Manager module - TBD diff --git a/ml/src/data_pipeline/mod.rs b/ml/src/data_pipeline/mod.rs new file mode 100644 index 000000000..09468c066 --- /dev/null +++ b/ml/src/data_pipeline/mod.rs @@ -0,0 +1,264 @@ +//! Data Pipeline for ML Training +//! +//! Automates data downloading, caching, and train/val/test splitting. +//! Provides `DatasetSpec` for declarative dataset definitions and +//! `DatasetManager` for orchestrating the full pipeline. + +pub mod cache; +pub mod manager; +pub mod prepared; + +use chrono::NaiveDate; +use serde::{Deserialize, Serialize}; + +/// Asset class categorization +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum AssetClass { + Equity, + ETF, + Future, +} + +/// Bar size / time resolution +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum BarSize { + OneMinute, + FiveMinute, + FifteenMinute, + OneHour, + Daily, +} + +impl BarSize { + /// Estimated bars per trading day (6.5 hours for equities) + pub fn bars_per_day(&self) -> usize { + match self { + BarSize::OneMinute => 390, + BarSize::FiveMinute => 78, + BarSize::FifteenMinute => 26, + BarSize::OneHour => 7, + BarSize::Daily => 1, + } + } +} + +/// Date range specification +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub enum DateRange { + /// Absolute date range + Absolute { start: NaiveDate, end: NaiveDate }, + /// Rolling window from today + Rolling { days: u32 }, +} + +impl DateRange { + /// Resolve to absolute start/end dates + pub fn resolve(&self) -> (NaiveDate, NaiveDate) { + match self { + DateRange::Absolute { start, end } => (*start, *end), + DateRange::Rolling { days } => { + let end = chrono::Utc::now().date_naive(); + let start = end - chrono::Duration::days(i64::from(*days)); + (start, end) + } + } + } + + /// Number of calendar days in range + pub fn calendar_days(&self) -> u32 { + let (start, end) = self.resolve(); + (end - start).num_days().unsigned_abs() as u32 + } +} + +/// Train/validation/test split ratios +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct SplitRatio { + pub train: f64, + pub val: f64, + pub test: f64, +} + +impl Default for SplitRatio { + fn default() -> Self { + Self { + train: 0.70, + val: 0.15, + test: 0.15, + } + } +} + +impl SplitRatio { + /// Validate that ratios sum to 1.0 + pub fn is_valid(&self) -> bool { + let sum = self.train + self.val + self.test; + (sum - 1.0).abs() < 1e-6 && self.train > 0.0 && self.val >= 0.0 && self.test >= 0.0 + } +} + +/// Preset dataset modes +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum DatasetMode { + /// Last 30 days — fast iteration (~40K bars/symbol at 1m) + Dev, + /// Last 6 months — regime diversity (~120K bars/symbol at 1m) + Backtest, + /// 2+ years — production training (~480K bars/symbol at 1m) + Full, +} + +impl DatasetMode { + /// Default rolling window days for this mode + pub fn default_days(&self) -> u32 { + match self { + DatasetMode::Dev => 30, + DatasetMode::Backtest => 180, + DatasetMode::Full => 730, + } + } + + /// Create a DateRange for this mode + pub fn to_date_range(&self) -> DateRange { + DateRange::Rolling { + days: self.default_days(), + } + } +} + +/// Specification for a single symbol in the dataset +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct SymbolSpec { + /// Trading symbol (e.g., "AAPL", "SPY") + pub symbol: String, + /// Exchange identifier (e.g., "XNAS", "XNYS") + pub exchange: String, + /// Asset classification + pub asset_class: AssetClass, +} + +/// Complete dataset specification — declarative definition of training data needs +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct DatasetSpec { + /// Human-readable name (e.g., "us-equities-dev") + pub name: String, + /// Symbols to include + pub symbols: Vec, + /// Date range for data + pub date_range: DateRange, + /// Bar time resolution + pub bar_size: BarSize, + /// Train/val/test split ratios + pub split_ratio: SplitRatio, + /// Warmup bars for technical indicators (default: 50) + pub warmup_bars: usize, +} + +impl DatasetSpec { + /// Create a spec from a preset mode + pub fn from_mode(name: &str, symbols: Vec, mode: DatasetMode) -> Self { + Self { + name: name.to_string(), + symbols, + date_range: mode.to_date_range(), + bar_size: BarSize::OneMinute, + split_ratio: SplitRatio::default(), + warmup_bars: 50, + } + } + + /// Estimate total bars across all symbols + pub fn estimated_total_bars(&self) -> usize { + let trading_days = (self.date_range.calendar_days() as f64 * 252.0 / 365.0) as usize; + let bars_per_symbol = trading_days * self.bar_size.bars_per_day(); + bars_per_symbol * self.symbols.len() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn sample_symbols() -> Vec { + vec![ + SymbolSpec { + symbol: "AAPL".into(), + exchange: "XNAS".into(), + asset_class: AssetClass::Equity, + }, + SymbolSpec { + symbol: "SPY".into(), + exchange: "XNYS".into(), + asset_class: AssetClass::ETF, + }, + ] + } + + #[test] + fn test_dataset_mode_days() { + assert_eq!(DatasetMode::Dev.default_days(), 30); + assert_eq!(DatasetMode::Backtest.default_days(), 180); + assert_eq!(DatasetMode::Full.default_days(), 730); + } + + #[test] + fn test_bar_size_bars_per_day() { + assert_eq!(BarSize::OneMinute.bars_per_day(), 390); + assert_eq!(BarSize::Daily.bars_per_day(), 1); + } + + #[test] + fn test_split_ratio_default_valid() { + let ratio = SplitRatio::default(); + assert!(ratio.is_valid()); + } + + #[test] + fn test_split_ratio_invalid() { + let ratio = SplitRatio { + train: 0.5, + val: 0.5, + test: 0.5, + }; + assert!(!ratio.is_valid()); + } + + #[test] + fn test_date_range_rolling_days() { + let range = DateRange::Rolling { days: 30 }; + assert_eq!(range.calendar_days(), 30); + } + + #[test] + fn test_dataset_spec_from_mode() { + let spec = DatasetSpec::from_mode("test", sample_symbols(), DatasetMode::Dev); + assert_eq!(spec.name, "test"); + assert_eq!(spec.symbols.len(), 2); + assert_eq!(spec.warmup_bars, 50); + assert!(spec.split_ratio.is_valid()); + } + + #[test] + fn test_estimated_total_bars() { + let spec = DatasetSpec::from_mode("test", sample_symbols(), DatasetMode::Dev); + let bars = spec.estimated_total_bars(); + // 30 days * ~252/365 trading ratio * 390 bars/day * 2 symbols + assert!(bars > 10_000); + assert!(bars < 100_000); + } + + #[test] + fn test_absolute_date_range() { + let range = DateRange::Absolute { + start: NaiveDate::from_ymd_opt(2024, 1, 1).unwrap_or_default(), + end: NaiveDate::from_ymd_opt(2024, 6, 30).unwrap_or_default(), + }; + assert_eq!(range.calendar_days(), 181); + } + + #[test] + fn test_asset_class_serde() { + let json = serde_json::to_string(&AssetClass::Equity).unwrap_or_default(); + assert!(json.contains("Equity")); + } +} diff --git a/ml/src/data_pipeline/prepared.rs b/ml/src/data_pipeline/prepared.rs new file mode 100644 index 000000000..1e58beb63 --- /dev/null +++ b/ml/src/data_pipeline/prepared.rs @@ -0,0 +1 @@ +//! Prepared module - TBD diff --git a/ml/src/lib.rs b/ml/src/lib.rs index 26ff71d7f..52ed0827d 100644 --- a/ml/src/lib.rs +++ b/ml/src/lib.rs @@ -832,6 +832,7 @@ pub mod real_data_loader; pub mod data_validation; pub mod random_model; pub mod model_registry; +pub mod data_pipeline; // ========== MISSING TYPES STUBS ==========