diff --git a/config/training/dqn-smoketest.toml b/config/training/dqn-smoketest.toml index ec7924ab0..3e5705ed5 100644 --- a/config/training/dqn-smoketest.toml +++ b/config/training/dqn-smoketest.toml @@ -21,6 +21,7 @@ hidden_dim_base = 64 reward_scale = 1.0 huber_delta = 1.0 symbol = "ES.FUT" +max_bars = 60000 imbalance_bar_threshold = 1.0 imbalance_bar_ewma_alpha = 0.1 diff --git a/crates/ml/src/trainers/dqn/config.rs b/crates/ml/src/trainers/dqn/config.rs index 022f310b2..99ac0858c 100644 --- a/crates/ml/src/trainers/dqn/config.rs +++ b/crates/ml/src/trainers/dqn/config.rs @@ -1218,6 +1218,10 @@ pub struct DQNHyperparameters { /// one instrument's DBN files are loaded. Default: "ES.FUT". pub symbol: String, + /// Maximum bars to load for training. 0 = unlimited. + /// Truncates dataset after loading, before walk-forward split. + pub max_bars: usize, + /// Imbalance bar threshold for MBP-10 to bar conversion. /// Higher = fewer bars (less noise), lower = more bars (more signal). /// Only used when `data_source = "mbp10"`. Default: 100.0. @@ -1705,6 +1709,7 @@ impl DQNHyperparameters { // Data source: "ohlcv" (backward compatible) or "mbp10" (imbalance bars) data_source: "ohlcv".to_string(), symbol: "ES.FUT".to_string(), + max_bars: 0, // 0 = unlimited imbalance_bar_threshold: 100.0, imbalance_bar_ewma_alpha: 0.1, diff --git a/crates/ml/src/trainers/dqn/data_loading.rs b/crates/ml/src/trainers/dqn/data_loading.rs index 2f3db094b..1b57fcb61 100644 --- a/crates/ml/src/trainers/dqn/data_loading.rs +++ b/crates/ml/src/trainers/dqn/data_loading.rs @@ -126,7 +126,7 @@ impl DQNTrainer { self.ofi_features = Some(Arc::from(cached.ofi)); } - let all_data: Vec<(FeatureVector, Vec)> = cached.features + let mut all_data: Vec<(FeatureVector, Vec)> = cached.features .into_iter() .zip(cached.targets.into_iter()) .map(|(f, t)| { @@ -134,6 +134,12 @@ impl DQNTrainer { }) .collect(); + // Truncate to max_bars if configured (0 = unlimited) + if self.hyperparams.max_bars > 0 && all_data.len() > self.hyperparams.max_bars { + tracing::info!("max_bars={}: truncating {} → {} bars", self.hyperparams.max_bars, all_data.len(), self.hyperparams.max_bars); + all_data.truncate(self.hyperparams.max_bars); + } + let split = (all_data.len() * 80) / 100; let train = all_data[..split].to_vec(); let val = all_data[split..].to_vec(); @@ -419,6 +425,12 @@ impl DQNTrainer { feature_dim ); + // Truncate to max_bars if configured (0 = unlimited) + if self.hyperparams.max_bars > 0 && training_data.len() > self.hyperparams.max_bars { + tracing::info!("max_bars={}: truncating {} → {} bars", self.hyperparams.max_bars, training_data.len(), self.hyperparams.max_bars); + training_data.truncate(self.hyperparams.max_bars); + } + // Split training data 80/20 for train/validation let split_idx = (training_data.len() * 80) / 100; let train_data = training_data[..split_idx].to_vec(); // cpu-side split diff --git a/crates/ml/src/training_profile.rs b/crates/ml/src/training_profile.rs index 3109a6c61..3298e8843 100644 --- a/crates/ml/src/training_profile.rs +++ b/crates/ml/src/training_profile.rs @@ -92,6 +92,10 @@ pub struct TrainingSection { /// EWMA alpha for adaptive imbalance threshold. Default: 0.1. pub imbalance_bar_ewma_alpha: Option, + + /// Maximum bars to load for training. 0 = unlimited (load all data). + /// Useful for fast smoketests on large datasets. + pub max_bars: Option, } /// Epsilon-greedy exploration parameters (DQN-specific). @@ -741,6 +745,9 @@ impl DqnTrainingProfile { if let Some(ref v) = t.symbol { hp.symbol = v.clone(); } + if let Some(v) = t.max_bars { + hp.max_bars = v; + } if let Some(v) = t.imbalance_bar_threshold { hp.imbalance_bar_threshold = v; }