From d99abd384458a0be35db1f3525194dc762637894 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 4 Apr 2026 09:54:43 +0200 Subject: [PATCH] feat: train(data_dir, symbol, cb) + strict fxcache discovery via shared function Co-Authored-By: Claude Opus 4.6 (1M context) --- crates/ml/src/trainers/dqn/data_loading.rs | 158 +++++------------- .../src/trainers/dqn/smoke_tests/helpers.rs | 2 +- crates/ml/src/trainers/dqn/trainer/mod.rs | 3 +- 3 files changed, 49 insertions(+), 114 deletions(-) diff --git a/crates/ml/src/trainers/dqn/data_loading.rs b/crates/ml/src/trainers/dqn/data_loading.rs index a3c004814..2f3db094b 100644 --- a/crates/ml/src/trainers/dqn/data_loading.rs +++ b/crates/ml/src/trainers/dqn/data_loading.rs @@ -93,129 +93,63 @@ impl DQNTrainer { pub async fn load_training_data( &mut self, dbn_data_dir: &str, + symbol: &str, ) -> Result<( Vec<(FeatureVector, Vec)>, Vec<(FeatureVector, Vec)>, )> { - // Symbol scoping: if symbol is set, load only that instrument's subdirectory. - // Cache key uses the base dir (matches precompute_features), data loading uses scoped dir. - let base_data_dir = dbn_data_dir; - let effective_data_dir = if !self.hyperparams.symbol.is_empty() { - let scoped = Path::new(base_data_dir).join(&self.hyperparams.symbol); - if scoped.exists() { - info!("Symbol filter: loading {} from {}", self.hyperparams.symbol, scoped.display()); - scoped.to_string_lossy().to_string() - } else { - base_data_dir.to_string() - } - } else { - base_data_dir.to_string() - }; - let dbn_data_dir = &effective_data_dir; - - // ── .fxcache (flat binary feature cache) — fastest path ────────────── - // Auto-discover cache dir: explicit > env var > sibling directory - let fxcache_dir = self.feature_cache_dir.clone().or_else(|| { - if let Ok(dir) = std::env::var("FOXHUNT_FEATURE_CACHE_DIR") { - let p = std::path::PathBuf::from(dir); - if p.exists() { return Some(p); } - } - // Walk up from data_dir to find sibling feature-cache/ - // e.g. test_data/futures-baseline/ES.FUT -> test_data/feature-cache - let mut dir = Path::new(dbn_data_dir); - loop { - if let Some(parent) = dir.parent() { - let candidate = parent.join("feature-cache"); - if candidate.exists() { return Some(candidate); } - if parent == dir { break; } // root - dir = parent; - } else { - break; - } - } + // ── fxcache discovery (strict key match via shared function) ──────── + let cache_dir_override = self.feature_cache_dir.as_deref(); + let mbp10_path = if self.hyperparams.mbp10_data_dir.is_empty() { None - }); - if let Some(ref cache_dir) = fxcache_dir { - // Cache key uses base dir (before symbol scoping) to match precompute_features - let data_dir_path = Path::new(base_data_dir); - // Resolve mbp10/trades paths: if relative and not found from CWD, - // try relative to data_dir ancestors (handles cargo test CWD = crates/ml/) - let resolve_sibling = |raw: &str| -> Option { - if raw.is_empty() { return None; } - let p = Path::new(raw); - if p.exists() { return Some(p.to_path_buf()); } - // Walk up from data_dir looking for the relative path - let mut dir = data_dir_path; - while let Some(parent) = dir.parent() { - let candidate = parent.join(raw); - if candidate.exists() { return Some(candidate); } - if parent == dir { break; } - dir = parent; - } - None - }; - let mbp10_dir = resolve_sibling(&self.hyperparams.mbp10_data_dir); - let trades_dir = resolve_sibling(&self.hyperparams.trades_data_dir); + } else { + let p = std::path::Path::new(&self.hyperparams.mbp10_data_dir); + if p.exists() { Some(p) } else { None } + }; + let trades_path = if self.hyperparams.trades_data_dir.is_empty() { + None + } else { + let p = std::path::Path::new(&self.hyperparams.trades_data_dir); + if p.exists() { Some(p) } else { None } + }; - let mut fxcache_path = crate::feature_cache::calculate_dbn_cache_key_full( - data_dir_path, mbp10_dir.as_deref(), trades_dir.as_deref(), - &self.hyperparams.symbol, &self.hyperparams.data_source, - ).ok() - .and_then(|hex| hex::decode(&hex).ok()) - .and_then(|b| <[u8; 32]>::try_from(b).ok()) - .and_then(|key| crate::fxcache::find_fxcache(cache_dir, &key)); - - if fxcache_path.is_none() { - // No exact key match — load the most recent .fxcache anyway. - // The cache was expensive to precompute; don't delete it just - // because the key changed (path or code changes shift the hash). - if let Ok(entries) = std::fs::read_dir(cache_dir) { - let mut candidates: Vec<_> = entries.flatten() - .filter(|e| e.path().extension().and_then(|s| s.to_str()) == Some("fxcache")) - .collect(); - candidates.sort_by_key(|e| std::cmp::Reverse(e.metadata().ok().and_then(|m| m.modified().ok()))); - if let Some(newest) = candidates.first() { - info!("fxcache key mismatch — using most recent cache: {:?}", newest.path()); - fxcache_path = Some(newest.path()); - } - } + if let Some(cached) = crate::fxcache::discover_and_load( + std::path::Path::new(dbn_data_dir), + symbol, + mbp10_path, + trades_path, + &self.hyperparams.data_source, + cache_dir_override, + ) { + // Set OFI from cache using explicit has_ofi flag (not zero-detection) + if cached.has_ofi { + self.ofi_features = Some(Arc::from(cached.ofi)); } - if let Some(fxcache_path) = fxcache_path { - match crate::fxcache::load_fxcache(&fxcache_path) { - Ok(cached) => { - info!( - "fxcache hit: {} bars from {:?}", - cached.bar_count, fxcache_path - ); + let all_data: Vec<(FeatureVector, Vec)> = cached.features + .into_iter() + .zip(cached.targets.into_iter()) + .map(|(f, t)| { + (f, t.to_vec()) + }) + .collect(); - // Set OFI features from cache - let has_ofi = cached.ofi.iter().any(|o| o.iter().any(|&v| v != 0.0)); - if has_ofi { - self.ofi_features = Some(Arc::from(cached.ofi)); - } - - // Combine features + targets - let all_data: Vec<([f64; 42], Vec)> = cached - .features - .into_iter() - .zip(cached.targets.into_iter()) - .map(|(f, t)| (f, t.to_vec())) - .collect(); - - // 80/20 split (same as other cache paths) - let split = (all_data.len() * 80) / 100; - let train = all_data[..split].to_vec(); - let val = all_data[split..].to_vec(); - return Ok((train, val)); - } - Err(e) => { - debug!("fxcache load failed: {e}, falling through"); - } - } - } + let split = (all_data.len() * 80) / 100; + let train = all_data[..split].to_vec(); + let val = all_data[split..].to_vec(); + return Ok((train, val)); } + // ── DBN fallback: load bars scoped to symbol ───────────────────────── + let symbol_data_dir = std::path::Path::new(dbn_data_dir).join(symbol); + let effective_dir = if symbol_data_dir.exists() { + info!("Symbol filter: loading {} from {}", symbol, symbol_data_dir.display()); + symbol_data_dir.to_string_lossy().to_string() + } else { + dbn_data_dir.to_string() + }; + let dbn_data_dir = &effective_dir; + // ── Load bars: MBP-10 imbalance bars or OHLCV candles ───────────── let all_ohlcv_bars = match self.hyperparams.data_source.as_str() { "mbp10" => { diff --git a/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs b/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs index b6cf419cb..28df328a4 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/helpers.rs @@ -121,7 +121,7 @@ pub(super) fn load_smoke_data() -> anyhow::Result<( let rt = tokio::runtime::Builder::new_current_thread() .enable_all() .build()?; - let (train, val) = rt.block_on(trainer.load_training_data(&data_dir))?; + let (train, val) = rt.block_on(trainer.load_training_data(&data_dir, "ES.FUT"))?; Ok((train, val)) } diff --git a/crates/ml/src/trainers/dqn/trainer/mod.rs b/crates/ml/src/trainers/dqn/trainer/mod.rs index 3fc01e8cc..f744806a8 100644 --- a/crates/ml/src/trainers/dqn/trainer/mod.rs +++ b/crates/ml/src/trainers/dqn/trainer/mod.rs @@ -431,6 +431,7 @@ impl DQNTrainer { pub async fn train( &mut self, dbn_data_dir: &str, + symbol: &str, checkpoint_callback: F, ) -> Result where @@ -442,7 +443,7 @@ impl DQNTrainer { ); // Load market data from DBN files (ALL data for walk-forward or single-pass) - let (training_data, val_data) = self.load_training_data(dbn_data_dir).await?; + let (training_data, val_data) = self.load_training_data(dbn_data_dir, symbol).await?; info!( "Loaded {} training samples, {} validation samples",