diff --git a/crates/ml/examples/precompute_features.rs b/crates/ml/examples/precompute_features.rs index 3d3796c58..be48e3b5a 100644 --- a/crates/ml/examples/precompute_features.rs +++ b/crates/ml/examples/precompute_features.rs @@ -96,7 +96,8 @@ use clap::Parser; use tracing::info; use ml::trainers::dqn::{ - collect_dbn_files_recursive, extract_features_from_bars, extract_ohlcv_bars_from_dbn, + collect_dbn_files_filtered, collect_dbn_files_recursive, extract_features_from_bars, + extract_ohlcv_bars_from_dbn, }; use ml::features::extraction::OHLCVBar; @@ -210,10 +211,24 @@ async fn main() -> Result<()> { let t0 = Instant::now(); // ── Step 1: Load OHLCV bars from DBN files ────────────────────────────── - info!("Loading OHLCV bars from {}...", data_dir.display()); - let mut dbn_files = collect_dbn_files_recursive(&data_dir); + info!("Loading OHLCV bars from {} (symbol={})...", data_dir.display(), opts.symbol); + let mut dbn_files = collect_dbn_files_filtered(&data_dir, Some(&opts.symbol)); + if dbn_files.is_empty() { + anyhow::bail!( + "No DBN files found for symbol '{}' in {}. Available: {:?}", + opts.symbol, + data_dir.display(), + std::fs::read_dir(&data_dir) + .ok() + .map(|d| d.filter_map(|e| e.ok()) + .filter(|e| e.path().is_dir()) + .map(|e| e.file_name().to_string_lossy().into_owned()) + .collect::>()) + .unwrap_or_default() + ); + } dbn_files.sort(); - info!("Found {} DBN files", dbn_files.len()); + info!("Found {} DBN files for symbol '{}'", dbn_files.len(), opts.symbol); let mut all_bars: Vec = Vec::new(); for file in &dbn_files { diff --git a/crates/ml/src/trainers/dqn/data_loading.rs b/crates/ml/src/trainers/dqn/data_loading.rs index 860a4f60d..23dcad274 100644 --- a/crates/ml/src/trainers/dqn/data_loading.rs +++ b/crates/ml/src/trainers/dqn/data_loading.rs @@ -48,6 +48,32 @@ pub fn collect_dbn_files_recursive(dir: &Path) -> Vec { files } +/// Collect DBN files, filtering to a specific symbol subdirectory. +/// If `symbol` is `Some("ES.FUT")`, only loads files from `dir/ES.FUT/`. +/// If the symbol subdirectory doesn't exist, falls back to filename matching. +pub fn collect_dbn_files_filtered(dir: &Path, symbol: Option<&str>) -> Vec { + match symbol { + Some(sym) => { + let symbol_dir = dir.join(sym); + if symbol_dir.is_dir() { + collect_dbn_files_recursive(&symbol_dir) + } else { + // No subdirectory — filter files by name containing symbol + collect_dbn_files_recursive(dir) + .into_iter() + .filter(|p| { + p.file_name() + .and_then(|n| n.to_str()) + .map(|n| n.contains(sym)) + .unwrap_or(false) + }) + .collect() + } + } + None => collect_dbn_files_recursive(dir), + } +} + /// Extract OHLCV bars from a single DBN file — no DQNTrainer required. pub fn extract_ohlcv_bars_from_dbn(file_path: &Path) -> Result> { if is_zstd_file(file_path)? { diff --git a/crates/ml/src/trainers/dqn/mod.rs b/crates/ml/src/trainers/dqn/mod.rs index ce1b50265..171b49517 100644 --- a/crates/ml/src/trainers/dqn/mod.rs +++ b/crates/ml/src/trainers/dqn/mod.rs @@ -38,7 +38,7 @@ pub use config::{DQNAgentType, DQNHyperparameters}; pub use early_stopping::EarlyStopping; pub use lr_scheduler::{LRDecayType, LRScheduler}; pub use statistics::{FeatureStatistics, QValueStats}; -pub use data_loading::{collect_dbn_files_recursive, extract_ohlcv_bars_from_dbn}; +pub use data_loading::{collect_dbn_files_filtered, collect_dbn_files_recursive, extract_ohlcv_bars_from_dbn}; pub use features::extract_features_from_bars; pub use trainer::DQNTrainer;