//! Walk-forward training binary for DQN and PPO models. //! //! Trains models using expanding walk-forward windows on real OHLCV data loaded //! from Databento DBN files. Supports early stopping, checkpoint saving, and //! normalization statistics export for reproducible evaluation. //! //! # Usage //! //! ```bash //! SQLX_OFFLINE=true cargo run -p ml --example train_baseline -- \ //! --model both --epochs 50 --batch-size 128 \ //! --data-dir data/cache/futures-baseline \ //! --output-dir ml/trained_models //! ``` #![allow(unused_crate_dependencies)] #![deny( clippy::unwrap_used, clippy::expect_used, clippy::panic, clippy::indexing_slicing )] use std::path::{Path, PathBuf}; use anyhow::{Context, Result}; use clap::Parser; use rand::Rng; use tracing::{error, info, warn}; use ml::dqn::{DQNConfig, Experience, DQN}; #[allow(unreachable_pub)] mod baseline_common; use baseline_common::{load_all_bars, spread_cost_bps}; use ml::features::extraction::extract_ml_features; use ml::ppo::ppo::{PPOConfig, PPO}; use ml::ppo::trajectories::{Trajectory, TrajectoryBatch, TrajectoryStep}; use ml::ppo::gae::compute_gae; use ml::types::OHLCVBar; use ml::walk_forward::{generate_walk_forward_windows, NormStats, WalkForwardConfig}; // --------------------------------------------------------------------------- // CLI Arguments // --------------------------------------------------------------------------- /// Walk-forward training binary for DQN and PPO baseline models. #[derive(Parser, Debug)] #[command(name = "train_baseline", about = "Train DQN/PPO with walk-forward windows")] struct Args { /// Which model(s) to train: "dqn", "ppo", or "both" #[arg(long, default_value = "both")] model: String, /// Maximum training epochs per fold #[arg(long, default_value_t = 50)] epochs: usize, /// Training batch size #[arg(long, default_value_t = 128)] batch_size: usize, /// Path to directory containing .dbn.zst files #[arg(long, default_value = "data/cache/futures-baseline")] data_dir: PathBuf, /// Output directory for trained model checkpoints #[arg(long, default_value = "ml/trained_models")] output_dir: PathBuf, /// Optional path to hyperopt results JSON (reserved for future use) #[arg(long)] hyperopt_params: Option, /// Feature dimension (must match extract_ml_features output) #[arg(long, default_value_t = 51)] feature_dim: usize, /// Early stopping patience (epochs without improvement) #[arg(long, default_value_t = 10)] patience: usize, /// Number of actions for the DQN/PPO action space #[arg(long, default_value_t = 3)] num_actions: usize, /// Walk-forward: initial training window in months #[arg(long, default_value_t = 12)] train_months: u32, /// Walk-forward: validation window in months #[arg(long, default_value_t = 3)] val_months: u32, /// Walk-forward: test window in months #[arg(long, default_value_t = 3)] test_months: u32, /// Walk-forward: step size in months between folds #[arg(long, default_value_t = 3)] step_months: u32, /// Learning rate for optimizer #[arg(long, default_value_t = 1e-4)] learning_rate: f64, /// Max environment steps per epoch (caps trajectory length for OOM safety; /// 0 = use all bars, but this can use >1GB RAM per fold on large datasets) #[arg(long, default_value_t = 2000)] max_steps_per_epoch: usize, /// Symbol subdirectory to load (e.g. "ES.FUT", "NQ.FUT") #[arg(long, default_value = "ES.FUT")] symbol: String, /// Maximum absolute per-bar return; larger moves are clamped (contract roll filter) #[arg(long, default_value_t = 0.01)] max_bar_return: f64, /// Round-trip commission cost in basis points (1 bps = 0.01%) /// Applied to BUY/SELL rewards; HOLD is free. /// Default 1.0 bps covers ~$4 exchange+broker for ES e-mini. #[arg(long, default_value_t = 1.0)] tx_cost_bps: f64, /// Instrument tick size in price units (ES=0.25, NQ=0.25, ZN=1/64) #[arg(long, default_value_t = 0.25)] tick_size: f64, /// Typical bid-ask spread in ticks (ES=1.0, ZN=1.0, 6E=2.0) /// Half-spread slippage is added to tx_cost_bps per trade. #[arg(long, default_value_t = 1.0)] spread_ticks: f64, } // --------------------------------------------------------------------------- // Reward Calculation // --------------------------------------------------------------------------- /// Compute return-based reward for a trading action (in basis points). /// /// Uses percentage returns scaled to basis points (1bp = 0.01%) for stable /// Q-value magnitudes. Total cost (commission + spread slippage) is subtracted /// from BUY/SELL rewards to teach the model that trading isn't free. /// /// - action 0 (BUY): reward = return_bps - total_cost_bps /// - action 1 (SELL): reward = -return_bps - total_cost_bps /// - action 2 (HOLD): reward = 0 (no cost, no return) fn compute_reward(close_current: f64, close_next: f64, action_idx: u8, total_cost_bps: f64) -> f32 { // Convert to basis points: (next - cur) / cur * 10_000, clipped to ±10 bps // to prevent outlier rewards from destabilizing Q-values. let return_bps = if close_current.abs() > 1e-10 { (close_next - close_current) / close_current * 10_000.0 } else { 0.0 }; let clipped = return_bps.clamp(-10.0, 10.0); let reward = match action_idx { 0 => clipped - total_cost_bps, // BUY: profit minus cost 1 => -clipped - total_cost_bps, // SELL: profit minus cost _ => 0.0, // HOLD: no position, no cost }; reward as f32 } // --------------------------------------------------------------------------- // DQN Training // --------------------------------------------------------------------------- /// Train a DQN model on a single walk-forward fold. /// /// Returns the best validation loss achieved. fn train_dqn_fold( fold: usize, train_features: &[[f64; 51]], val_features: &[[f64; 51]], train_bars: &[OHLCVBar], val_bars: &[OHLCVBar], args: &Args, output_dir: &Path, ) -> Result { // Caller must pre-align bars to features (skip warmup period before calling). // If these differ, bar-index reward computation will silently use wrong prices. debug_assert_eq!( train_bars.len(), train_features.len(), "train_bars ({}) must be pre-aligned to train_features ({})", train_bars.len(), train_features.len(), ); debug_assert_eq!( val_bars.len(), val_features.len(), "val_bars ({}) must be pre-aligned to val_features ({})", val_bars.len(), val_features.len(), ); info!(" [DQN] Fold {} — {} train, {} val features", fold, train_features.len(), val_features.len()); // Configure DQN with baseline-friendly settings let config = DQNConfig { state_dim: args.feature_dim, num_actions: args.num_actions, hidden_dims: vec![128, 64], learning_rate: args.learning_rate, gamma: 0.95, // shorter horizon (20 bars vs 100) epsilon_start: 1.0, epsilon_end: 0.05, epsilon_decay: 0.995, replay_buffer_capacity: 50_000, batch_size: args.batch_size, min_replay_size: args.batch_size, target_update_freq: 200, // more frequent target sync warmup_steps: 0, // No warmup — we fill buffer before training use_double_dqn: true, use_huber_loss: true, // Disable Rainbow extras for baseline simplicity use_per: false, use_dueling: false, use_distributional: false, use_noisy_nets: false, use_cql: false, use_iqn: false, use_cvar_action_selection: false, ..DQNConfig::default() }; let mut dqn = DQN::new(config).context("Failed to create DQN model")?; let mut best_val_loss = f64::MAX; let mut epochs_without_improvement = 0_usize; let mut rng = rand::thread_rng(); for epoch in 0..args.epochs { // --- Training pass --- let mut epoch_loss = 0.0_f64; let mut epoch_steps = 0_usize; // Process training data sequentially let n_train = train_features.len(); if n_train < 2 { warn!(" [DQN] Fold {} — insufficient training features ({})", fold, n_train); break; } let step_limit = if args.max_steps_per_epoch > 0 { args.max_steps_per_epoch.min(n_train.saturating_sub(1)) } else { n_train.saturating_sub(1) }; // Features skip warmup period, so feature[i] corresponds to // train_bars[i + warmup_offset]. We must use the same offset for // reward computation to avoid misalignment. let warmup_offset = train_bars.len().saturating_sub(n_train); for i in 0..step_limit { let state_f64 = match train_features.get(i) { Some(f) => f, None => continue, }; let next_f64 = match train_features.get(i + 1) { Some(f) => f, None => continue, }; let state: Vec = state_f64.iter().map(|&v| v as f32).collect(); let next_state: Vec = next_f64.iter().map(|&v| v as f32).collect(); // Select action with epsilon-greedy let action = dqn.select_action(&state) .map(|fa| fa.to_index().min(args.num_actions.saturating_sub(1)) as u8) .unwrap_or_else(|_| rng.gen_range(0..args.num_actions as u8)); // Compute reward from bar prices (offset by warmup to align with features) let bar_idx = i + warmup_offset; let close_cur = train_bars.get(bar_idx).map(|b| b.close).unwrap_or(0.0); let close_next = train_bars.get(bar_idx + 1).map(|b| b.close).unwrap_or(close_cur); let total_cost = args.tx_cost_bps + spread_cost_bps(close_cur, args.tick_size, args.spread_ticks); let reward = compute_reward(close_cur, close_next, action, total_cost); let done = i + 2 >= n_train; // Store experience let exp = Experience::new(state, action, reward, next_state, done); if let Err(e) = dqn.store_experience(exp) { // Non-fatal: buffer may not be ready if epoch == 0 && i < 5 { info!(" [DQN] store_experience: {}", e); } } // Train step (returns (loss, grad_norm)) match dqn.train_step(None) { Ok((loss, _grad_norm)) => { epoch_loss += loss as f64; epoch_steps += 1; } Err(_) => { // Training not ready yet (buffer too small) } } } let avg_train_loss = if epoch_steps > 0 { epoch_loss / epoch_steps as f64 } else { f64::MAX }; // --- Validation pass --- let val_loss = evaluate_dqn_validation(&mut dqn, val_features, val_bars, args); // Decay epsilon at epoch level let current_eps = dqn.get_epsilon(); let new_eps = (current_eps * 0.995_f32).max(0.05); dqn.set_epsilon(new_eps as f64); info!( " [DQN] Fold {} Epoch {}/{} — train_loss={:.6} val_loss={:.6} eps={:.4}", fold, epoch + 1, args.epochs, avg_train_loss, val_loss, dqn.get_epsilon() ); // Early stopping check if val_loss < best_val_loss { best_val_loss = val_loss; epochs_without_improvement = 0; // Save best checkpoint let ckpt_path = output_dir.join(format!("dqn_fold{}_best.safetensors", fold)); if let Err(e) = dqn.get_q_network_vars().save(ckpt_path.to_string_lossy().as_ref()) { warn!(" [DQN] Failed to save checkpoint: {}", e); } else { info!(" [DQN] Saved best checkpoint: {}", ckpt_path.display()); } } else { epochs_without_improvement += 1; if epochs_without_improvement >= args.patience { info!( " [DQN] Early stopping at epoch {} (patience {} exhausted)", epoch + 1, args.patience ); break; } } } Ok(best_val_loss) } /// Evaluate DQN on validation features and return average loss proxy. /// /// Since DQN.train_step uses replay buffer internally, we estimate validation /// performance via average absolute reward (lower is closer to zero = better). fn evaluate_dqn_validation( dqn: &mut DQN, val_features: &[[f64; 51]], val_bars: &[OHLCVBar], args: &Args, ) -> f64 { // Bars must be pre-aligned to features by the caller. debug_assert_eq!( val_bars.len(), val_features.len(), "val_bars ({}) must be pre-aligned to val_features ({})", val_bars.len(), val_features.len(), ); // Run greedy inference on validation set and compute negative average reward. // Lower values = better model (early stopping minimizes this). let n = val_features.len(); if n < 2 { return f64::MAX; } // Save epsilon and force greedy for validation. // Note: select_action increments total_steps (side effect), but this is benign // because warmup_steps=0 in our config, so the warmup branch is never taken. let saved_epsilon = dqn.get_epsilon(); dqn.set_epsilon(0.0); let warmup_offset = val_bars.len().saturating_sub(n); let mut total_reward = 0.0_f64; let mut count = 0_usize; for i in 0..n.saturating_sub(1) { let state: Vec = match val_features.get(i) { Some(f) => f.iter().map(|&v| v as f32).collect(), None => break, }; // Greedy action (epsilon=0) let action = match dqn.select_action(&state) { Ok(fa) => fa.to_index().min(2) as u8, Err(_) => 2, // HOLD on error }; let bar_idx = i + warmup_offset; let close_cur = val_bars.get(bar_idx).map(|b| b.close).unwrap_or(0.0); let close_next = val_bars.get(bar_idx + 1).map(|b| b.close).unwrap_or(close_cur); let total_cost = args.tx_cost_bps + spread_cost_bps(close_cur, args.tick_size, args.spread_ticks); let reward = compute_reward(close_cur, close_next, action, total_cost); total_reward += reward as f64; count += 1; } // Restore epsilon dqn.set_epsilon(saved_epsilon as f64); if count > 0 { // Return negative mean reward: lower = better model -(total_reward / count as f64) } else { f64::MAX } } // --------------------------------------------------------------------------- // PPO Training // --------------------------------------------------------------------------- /// Train a PPO model on a single walk-forward fold. /// /// Returns the best validation loss achieved. fn train_ppo_fold( fold: usize, train_features: &[[f64; 51]], val_features: &[[f64; 51]], train_bars: &[OHLCVBar], val_bars: &[OHLCVBar], args: &Args, output_dir: &Path, ) -> Result { info!(" [PPO] Fold {} — {} train, {} val features", fold, train_features.len(), val_features.len()); let config = PPOConfig { state_dim: args.feature_dim, num_actions: args.num_actions, policy_hidden_dims: vec![128, 64], value_hidden_dims: vec![128, 64], policy_learning_rate: args.learning_rate, value_learning_rate: args.learning_rate * 3.0, clip_epsilon: 0.2, value_loss_coeff: 0.5, entropy_coeff: 0.01, batch_size: args.batch_size.max(64), mini_batch_size: 64, num_epochs: 4, max_grad_norm: 0.5, use_lstm: false, ..PPOConfig::default() }; let mut ppo = PPO::new(config.clone()).context("Failed to create PPO model")?; let mut best_val_loss = f64::MAX; let mut epochs_without_improvement = 0_usize; let n_train = train_features.len(); if n_train < 2 { warn!(" [PPO] Fold {} — insufficient training features ({})", fold, n_train); return Ok(f64::MAX); } for epoch in 0..args.epochs { // --- Collect trajectory --- let trajectory = collect_ppo_trajectory( &ppo, train_features, train_bars, args, )?; if trajectory.length < 2 { warn!(" [PPO] Fold {} Epoch {} — trajectory too short", fold, epoch + 1); continue; } // Compute GAE advantages and returns let trajectories = vec![trajectory]; let (advantages, returns) = match compute_gae(&trajectories, &config.gae_config) { Ok((adv, ret)) => (adv, ret), Err(e) => { warn!(" [PPO] Fold {} Epoch {} GAE failed: {}", fold, epoch + 1, e); continue; } }; let mut batch = TrajectoryBatch::from_trajectories( trajectories, advantages, returns, ); // PPO update match ppo.update(&mut batch) { Ok((policy_loss, value_loss)) => { // --- Validation pass --- let val_loss = evaluate_ppo_validation(&ppo, val_features, val_bars, args.tx_cost_bps, args.tick_size, args.spread_ticks); info!( " [PPO] Fold {} Epoch {}/{} — policy_loss={:.6} value_loss={:.6} val_metric={:.6}", fold, epoch + 1, args.epochs, policy_loss, value_loss, val_loss ); // Early stopping check if val_loss < best_val_loss { best_val_loss = val_loss; epochs_without_improvement = 0; // Save checkpoint let actor_path = output_dir.join(format!("ppo_fold{}_actor.safetensors", fold)); let critic_path = output_dir.join(format!("ppo_fold{}_critic.safetensors", fold)); let meta_path = output_dir.join(format!("ppo_fold{}_meta.json", fold)); if let Err(e) = ppo.save_checkpoint( &actor_path.to_string_lossy(), &critic_path.to_string_lossy(), &meta_path.to_string_lossy(), ) { warn!(" [PPO] Failed to save checkpoint: {}", e); } else { info!(" [PPO] Saved best checkpoint: {}", actor_path.display()); } } else { epochs_without_improvement += 1; if epochs_without_improvement >= args.patience { info!( " [PPO] Early stopping at epoch {} (patience {} exhausted)", epoch + 1, args.patience ); break; } } } Err(e) => { warn!(" [PPO] Fold {} Epoch {} update error: {}", fold, epoch + 1, e); } } } Ok(best_val_loss) } /// Collect a single trajectory from training data for PPO. /// /// Respects `--max-steps-per-epoch` to cap trajectory length and prevent OOM. /// Uses `act_with_log_prob()` to get real policy log-probabilities for /// importance sampling (instead of a uniform prior). fn collect_ppo_trajectory( ppo: &PPO, features: &[[f64; 51]], bars: &[OHLCVBar], args: &Args, ) -> Result { use ml::dqn::TradingAction; // Bars must be pre-aligned to features by the caller. debug_assert_eq!( bars.len(), features.len(), "bars ({}) must be pre-aligned to features ({})", bars.len(), features.len(), ); let mut trajectory = Trajectory::new(); let n = features.len(); let mut rng = rand::thread_rng(); // Cap steps to avoid OOM (same as DQN path) let step_limit = if args.max_steps_per_epoch > 0 { args.max_steps_per_epoch.min(n.saturating_sub(1)) } else { n.saturating_sub(1) }; // Warmup offset: features skip warmup period, align bar index let warmup_offset = bars.len().saturating_sub(n); for i in 0..step_limit { let state_f64 = match features.get(i) { Some(f) => f, None => break, }; let state: Vec = state_f64.iter().map(|&v| v as f32).collect(); // Get action, log_prob, and value from PPO policy let (action, log_prob, value) = match ppo.act_with_log_prob(&state) { Ok((a, lp, v)) => (a, lp, v), Err(_) => { // Fallback: random action with uniform log_prob and zero value let a = match rng.gen_range(0..args.num_actions) { 0 => TradingAction::Buy, 1 => TradingAction::Sell, _ => TradingAction::Hold, }; (a, -(args.num_actions as f32).ln(), 0.0_f32) } }; // Compute reward (offset by warmup to align features with bars) let bar_idx = i + warmup_offset; let close_cur = bars.get(bar_idx).map(|b| b.close).unwrap_or(0.0); let close_next = bars.get(bar_idx + 1).map(|b| b.close).unwrap_or(close_cur); let action_idx = action.to_int(); let total_cost = args.tx_cost_bps + spread_cost_bps(close_cur, args.tick_size, args.spread_ticks); let reward = compute_reward(close_cur, close_next, action_idx, total_cost); let done = i + 1 >= step_limit; let step = TrajectoryStep::new(state, action, log_prob, value, reward, done); trajectory.add_step(step); } Ok(trajectory) } /// Evaluate PPO validation performance (average absolute price change). fn evaluate_ppo_validation( ppo: &PPO, val_features: &[[f64; 51]], val_bars: &[OHLCVBar], tx_cost_bps: f64, tick_size: f64, spread_ticks_val: f64, ) -> f64 { // Bars must be pre-aligned to features by the caller. debug_assert_eq!( val_bars.len(), val_features.len(), "val_bars ({}) must be pre-aligned to val_features ({})", val_bars.len(), val_features.len(), ); // Run PPO inference on validation set and compute negative average reward. // Lower values = better model (early stopping minimizes this). let n = val_features.len(); if n < 2 { return f64::MAX; } let warmup_offset = val_bars.len().saturating_sub(n); let mut total_reward = 0.0_f64; let mut count = 0_usize; for i in 0..n.saturating_sub(1) { let state: Vec = match val_features.get(i) { Some(f) => f.iter().map(|&v| v as f32).collect(), None => break, }; // Greedy (deterministic) action for validation — avoids stochastic // noise that would make the early-stopping signal unreliable. let action_idx = match ppo.greedy_action(&state) { Ok(action) => action.to_int(), Err(_) => 2, // HOLD on error }; let bar_idx = i + warmup_offset; let close_cur = val_bars.get(bar_idx).map(|b| b.close).unwrap_or(0.0); let close_next = val_bars.get(bar_idx + 1).map(|b| b.close).unwrap_or(close_cur); let total_cost = tx_cost_bps + spread_cost_bps(close_cur, tick_size, spread_ticks_val); let reward = compute_reward(close_cur, close_next, action_idx, total_cost); total_reward += reward as f64; count += 1; } if count > 0 { -(total_reward / count as f64) } else { f64::MAX } } // --------------------------------------------------------------------------- // Main // --------------------------------------------------------------------------- fn main() -> Result<()> { // Initialize tracing tracing_subscriber::fmt() .with_env_filter( tracing_subscriber::EnvFilter::try_from_default_env() .unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")), ) .init(); let args = Args::parse(); let train_dqn = args.model == "dqn" || args.model == "both"; let train_ppo = args.model == "ppo" || args.model == "both"; info!("=== Walk-Forward Baseline Training ==="); info!(" Model(s): {}", args.model); info!(" Symbol: {}", args.symbol); info!(" Epochs: {}", args.epochs); info!(" Batch size: {}", args.batch_size); info!(" Data dir: {}", args.data_dir.display()); info!(" Output dir: {}", args.output_dir.display()); info!(" Feature dim: {}", args.feature_dim); info!(" Num actions: {}", args.num_actions); info!(" Learning rate: {:.1e}", args.learning_rate); info!(" Tx cost: {:.1} bps commission + {:.1} tick spread (tick_size={:.4})", args.tx_cost_bps, args.spread_ticks, args.tick_size); info!(" Patience: {}", args.patience); // 1. Load all OHLCV bars from DBN files info!("Step 1/5: Loading OHLCV bars from DBN files..."); let bars = load_all_bars(&args.data_dir, &args.symbol)?; if bars.is_empty() { anyhow::bail!("No bars loaded from {}", args.data_dir.display()); } info!(" Loaded {} bars ({} to {})", bars.len(), bars.first().map(|b| b.timestamp.to_string()).unwrap_or_default(), bars.last().map(|b| b.timestamp.to_string()).unwrap_or_default(), ); // 2. Extract features info!("Step 2/5: Extracting {}-dimensional features...", args.feature_dim); let all_features = extract_ml_features(&bars) .context("Feature extraction failed")?; info!(" Extracted {} feature vectors (warmup period consumed {} bars)", all_features.len(), bars.len().saturating_sub(all_features.len()), ); // Since features skip the warmup period, we need bars aligned to features. // Features start at bar index warmup_offset (typically 50). let warmup_offset = bars.len().saturating_sub(all_features.len()); let aligned_bars = if warmup_offset < bars.len() { &bars[warmup_offset..] } else { &bars }; // 3. Generate walk-forward windows info!("Step 3/5: Generating walk-forward windows..."); let wf_config = WalkForwardConfig { initial_train_months: args.train_months, val_months: args.val_months, test_months: args.test_months, step_months: args.step_months, }; let windows = generate_walk_forward_windows(aligned_bars, &wf_config); if windows.is_empty() { anyhow::bail!( "No walk-forward windows generated. Need at least {} months of data.", wf_config.initial_train_months + wf_config.val_months + wf_config.test_months ); } info!(" Generated {} walk-forward folds", windows.len()); // Create output directory std::fs::create_dir_all(&args.output_dir) .with_context(|| format!("Failed to create output dir: {}", args.output_dir.display()))?; // 4. Train each fold info!("Step 4/5: Training models on each fold..."); let mut dqn_results: Vec<(usize, f64)> = Vec::new(); let mut ppo_results: Vec<(usize, f64)> = Vec::new(); for window in &windows { info!("--- Fold {} ---", window.fold); info!( " Train: {} bars (up to {}), Val: {} bars (up to {}), Test: {} bars (up to {})", window.train.len(), window.train_end, window.val.len(), window.val_end, window.test.len(), window.test_end, ); // Extract features for this fold's train and val sets let train_feat = match extract_ml_features(&window.train) { Ok(f) => f, Err(e) => { warn!(" Fold {} — train feature extraction failed: {}", window.fold, e); continue; } }; let val_feat = match extract_ml_features(&window.val) { Ok(f) => f, Err(e) => { warn!(" Fold {} — val feature extraction failed: {}", window.fold, e); continue; } }; if train_feat.is_empty() || val_feat.is_empty() { warn!(" Fold {} — empty features, skipping", window.fold); continue; } // Compute NormStats from training features only let norm_stats = NormStats::from_features(&train_feat); // Normalize features let train_norm = norm_stats.normalize_batch(&train_feat); let val_norm = norm_stats.normalize_batch(&val_feat); // Save NormStats let norm_path = args.output_dir.join(format!("norm_stats_fold{}.json", window.fold)); let norm_json = serde_json::to_string_pretty(&norm_stats) .context("Failed to serialize NormStats")?; std::fs::write(&norm_path, norm_json) .with_context(|| format!("Failed to write {}", norm_path.display()))?; info!(" Saved NormStats to {}", norm_path.display()); // Aligned bars for features (train features skip warmup period of train bars) let train_warmup = window.train.len().saturating_sub(train_norm.len()); let train_bars_aligned = if train_warmup < window.train.len() { &window.train[train_warmup..] } else { &window.train }; let val_warmup = window.val.len().saturating_sub(val_norm.len()); let val_bars_aligned = if val_warmup < window.val.len() { &window.val[val_warmup..] } else { &window.val }; // Train DQN if train_dqn { match train_dqn_fold( window.fold, &train_norm, &val_norm, train_bars_aligned, val_bars_aligned, &args, &args.output_dir, ) { Ok(best_loss) => { dqn_results.push((window.fold, best_loss)); } Err(e) => { error!(" [DQN] Fold {} failed: {}", window.fold, e); } } } // Train PPO if train_ppo { match train_ppo_fold( window.fold, &train_norm, &val_norm, train_bars_aligned, val_bars_aligned, &args, &args.output_dir, ) { Ok(best_loss) => { ppo_results.push((window.fold, best_loss)); } Err(e) => { error!(" [PPO] Fold {} failed: {}", window.fold, e); } } } } // 5. Summary info!("Step 5/5: Training Summary"); info!(" ==================================="); if train_dqn { info!(" DQN Results ({} folds):", dqn_results.len()); for (fold, loss) in &dqn_results { info!(" Fold {}: best_val_metric = {:.6}", fold, loss); } if !dqn_results.is_empty() { let avg: f64 = dqn_results.iter().map(|(_, l)| l).sum::() / dqn_results.len() as f64; info!(" Average: {:.6}", avg); } } if train_ppo { info!(" PPO Results ({} folds):", ppo_results.len()); for (fold, loss) in &ppo_results { info!(" Fold {}: best_val_metric = {:.6}", fold, loss); } if !ppo_results.is_empty() { let avg: f64 = ppo_results.iter().map(|(_, l)| l).sum::() / ppo_results.len() as f64; info!(" Average: {:.6}", avg); } } info!(" Checkpoints saved to: {}", args.output_dir.display()); info!(" ==================================="); Ok(()) }