//! DQN Training Example //! //! Trains a DQN model on market data and saves checkpoints to disk. //! //! # Usage //! //! ```bash //! # Train with default parameters (100 epochs) //! cargo run -p ml --example train_dqn --release --features cuda //! //! # Custom epochs and output path //! cargo run -p ml --example train_dqn --release --features cuda -- \ //! --epochs 500 \ //! --output ml/trained_models/dqn_model.safetensors //! //! # Custom data directory //! cargo run -p ml --example train_dqn --release --features cuda -- \ //! --data-dir test_data/real/databento/ml_training \ //! --epochs 500 //! ``` // Use mimalloc allocator for 10-25% performance improvement #[cfg(feature = "mimalloc-allocator")] use mimalloc::MiMalloc; #[cfg(feature = "mimalloc-allocator")] #[global_allocator] static GLOBAL: MiMalloc = MiMalloc; use anyhow::{Context, Result}; use clap::Parser; use std::path::PathBuf; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; use tokio::signal; use tracing::{info, warn}; use tracing_subscriber::FmtSubscriber; use ml::checkpoint::{CheckpointConfig, CheckpointManager}; use ml::data_loaders::BarSamplingMethod; use ml::trainers::dqn::{DQNHyperparameters, DQNTrainer}; use ml::trainers::TargetUpdateMode; /// Train DQN model on market data #[derive(Debug, Parser)] #[command(name = "train_dqn", about = "Train DQN model on market data")] struct Opts { /// Number of training epochs #[arg(long, default_value = "100")] epochs: usize, /// Learning rate /// BUG #18 FIX (Wave 16S-V17): Reduced 10Ɨ to 0.00001 to prevent Q-value explosion /// after Bug #16 fix increased reward scale (raw portfolio values vs normalized). /// Previous: 0.0001 caused gradient collapse (Q-values hit 1000.0 clamp, grad_norm → 0). #[arg(long, default_value = "0.00001")] learning_rate: f64, /// Batch size (max 230 for RTX 3050 Ti 4GB) /// Optimal value from hyperopt (42 trials, 2025-10-31): 32 #[arg(long, default_value = "32")] batch_size: usize, /// Discount factor (gamma) /// Optimal value from hyperopt (42 trials, 2025-10-31): 0.9626 #[arg(long, default_value = "0.9626")] gamma: f64, /// Checkpoint save frequency (epochs) #[arg(long, default_value = "10")] checkpoint_frequency: usize, /// Output directory for trained model #[arg(long, default_value = "ml/trained_models")] output_dir: String, /// Data directory containing DBN files #[arg(long, default_value = "test_data/real/databento/ml_training")] data_dir: String, /// Parquet file path (overrides data_dir if specified) #[arg(long)] parquet_file: Option, /// Verbose logging #[arg(short, long)] verbose: bool, /// Enable early stopping (recommended, use --no-early-stopping to disable) #[arg(long)] early_stopping: bool, /// Disable early stopping #[arg(long)] no_early_stopping: bool, /// Q-value floor threshold for early stopping #[arg(long, default_value = "0.5")] q_value_floor: f64, /// Minimum loss improvement percentage for plateau detection #[arg(long, default_value = "2.0")] min_loss_improvement: f64, /// Plateau detection window size (epochs) /// Optimal value from hyperopt: 5 #[arg(long, default_value = "5")] plateau_window: usize, /// Minimum epochs before early stopping can trigger /// Updated to 50 to prevent premature stopping (was 10) #[arg(long, default_value = "50")] min_epochs_before_stopping: usize, /// Initial exploration rate (epsilon start) /// Updated to 0.3 for more initial exploration (was 1.0) #[arg(long, default_value = "0.3")] epsilon_start: f64, /// Final exploration rate (epsilon end) /// Updated to 0.05 to maintain exploration (was 0.01) #[arg(long, default_value = "0.05")] epsilon_end: f64, /// Exploration decay rate /// Updated to 0.995 for slower decay (was 0.9968) #[arg(long, default_value = "0.995")] epsilon_decay: f64, /// Replay buffer capacity /// Optimal value from hyperopt: 104346 #[arg(long, default_value = "104346")] buffer_size: usize, /// Minimum replay buffer size before training starts /// Updated to 500 for more diverse experiences (was auto-calculated as batch_size * 2 = 64) #[arg(long, default_value = "500")] min_replay_size: usize, /// Checkpoint directory (overrides output_dir for checkpoints) #[arg(long)] checkpoint_dir: Option, /// Alternative bar sampling method (time, tick, volume, dollar, imbalance, run) #[arg(long, default_value = "time")] bar_method: String, /// HOLD penalty weight (higher = stronger penalty for holding) #[arg(long, default_value = "0.01")] hold_penalty_weight: f64, /// Price movement threshold for dynamic HOLD reward (as fraction, e.g., 0.01 = 1%) #[arg(long, default_value = "0.01")] bar_threshold: Option, /// Disable preprocessing (use raw non-stationary prices - NOT RECOMMENDED) #[arg(long)] no_preprocessing: bool, /// Preprocessing window size (default: 50 bars) #[arg(long, default_value = "50")] preprocessing_window: i64, /// Preprocessing clip sigma (default: 5.0σ) #[arg(long, default_value = "5.0")] preprocessing_clip_sigma: f64, /// Warmup steps for random exploration before training (Rainbow DQN: 80K) /// Default: ADAPTIVE (0 for <200K steps, scaled 200K-1M, 80K for >1M) /// Set explicitly to override adaptive behavior #[arg(long)] warmup_steps: Option, /// Initial capital for portfolio trading (default: $100,000, minimum: $1,000) #[arg(long, default_value = "100000.0")] initial_capital: f32, /// Cash reserve requirement as a percentage of portfolio value (0.0-100.0) #[arg(long, default_value = "0.0")] cash_reserve_percent: f64, /// Polyak averaging coefficient (tau) for soft target updates (default: 1.0 = hard updates) /// Set to 0.001 for soft updates (Rainbow DQN: 693-step convergence half-life) /// Lower values = slower convergence, higher values = faster convergence #[arg(long, default_value = "1.0")] tau: f64, /// Use soft target updates (Polyak averaging) instead of hard updates /// Soft updates blend target network gradually with main network /// Default: hard updates (tau=1.0, complete replacement every N steps) #[arg(long)] soft_updates: bool, } #[tokio::main] async fn main() -> Result<()> { // Parse CLI options let opts = Opts::parse(); // Setup logging let level = if opts.verbose { tracing::Level::DEBUG } else { tracing::Level::INFO }; let subscriber = FmtSubscriber::builder().with_max_level(level).finish(); tracing::subscriber::set_global_default(subscriber) .context("Failed to set tracing subscriber")?; #[cfg(feature = "mimalloc-allocator")] info!("šŸš€ Using mimalloc allocator for improved performance"); #[cfg(not(feature = "mimalloc-allocator"))] info!("ā„¹ļø Using system allocator (consider --features mimalloc-allocator for 10-25% speedup)"); info!("šŸš€ Starting DQN Training"); info!("Configuration:"); info!(" • Epochs: {}", opts.epochs); info!(" • Learning rate: {}", opts.learning_rate); info!(" • Batch size: {}", opts.batch_size); info!(" • Gamma: {}", opts.gamma); info!( " • Checkpoint frequency: {} epochs", opts.checkpoint_frequency ); info!(" • Output directory: {}", opts.output_dir); info!(" • Data directory: {}", opts.data_dir); info!(" • Bar sampling method: {}", opts.bar_method); if let Some(threshold) = opts.bar_threshold { info!(" • Bar threshold: {}", threshold); } info!(" • Epsilon start: {}", opts.epsilon_start); info!(" • Epsilon end: {}", opts.epsilon_end); info!(" • Epsilon decay: {}", opts.epsilon_decay); info!(" • Buffer size: {}", opts.buffer_size); info!(" • Min replay size: {}", opts.min_replay_size); info!(" • Initial capital: ${:.2}", opts.initial_capital); info!(" • Cash reserve: {}%", opts.cash_reserve_percent); // Log target update configuration if opts.soft_updates { info!(" • Target update mode: Soft (Polyak averaging)"); info!( " • Tau (Ļ„): {} (convergence half-life: {:.0} steps)", opts.tau, (-0.5_f64.ln()) / (-(1.0 - opts.tau).ln()) ); } else { info!(" • Target update mode: Hard (complete replacement every 10K steps)"); info!(" • Tau (Ļ„): {} (no blending, hard copy)", opts.tau); } // ═══════════════════════════════════════════════════════════════════════════ // ADAPTIVE WARMUP CALCULATION // ═══════════════════════════════════════════════════════════════════════════ // Estimate total training steps for adaptive warmup let avg_steps_per_epoch = 1392; // Empirical: ES_FUT_180d.parquet has 1,392 steps/epoch let total_steps_estimate = opts.epochs * avg_steps_per_epoch; info!( "šŸ“Š Training length estimate: {}K steps ({} epochs Ɨ {} steps/epoch)", total_steps_estimate / 1000, opts.epochs, avg_steps_per_epoch ); // Adaptive warmup calculation let effective_warmup = if let Some(explicit_warmup) = opts.warmup_steps { // User explicitly set warmup - respect it info!( "šŸŽÆ Using EXPLICIT warmup: {}K steps (user override)", explicit_warmup / 1000 ); explicit_warmup } else { // Adaptive warmup based on training length let adaptive_warmup = match total_steps_estimate { 0..=200_000 => { info!("⚔ ADAPTIVE WARMUP: 0 steps (short training <200K steps)"); 0 }, 200_001..=500_000 => { let warmup = total_steps_estimate / 20; // 5% warmup info!( "⚔ ADAPTIVE WARMUP: {}K steps (~5% of {}K total)", warmup / 1000, total_steps_estimate / 1000 ); warmup }, 500_001..=1_000_000 => { let warmup = total_steps_estimate / 12; // ~8% warmup info!( "⚔ ADAPTIVE WARMUP: {}K steps (~8% of {}K total)", warmup / 1000, total_steps_estimate / 1000 ); warmup }, _ => { info!("⚔ ADAPTIVE WARMUP: 80K steps (Rainbow DQN standard for >1M steps)"); 80_000 // Full Rainbow warmup for very long runs }, }; adaptive_warmup }; // Warning: warmup consuming too much of training if effective_warmup > 0 { let warmup_ratio = effective_warmup as f64 / total_steps_estimate as f64; if warmup_ratio > 0.10 { warn!( "āš ļø Warmup period ({}K steps) is {:.1}% of total training ({}K steps)", effective_warmup / 1000, warmup_ratio * 100.0, total_steps_estimate / 1000 ); warn!("āš ļø This may significantly delay learning. Consider:"); warn!( " • Increase --epochs to lengthen training (recommended: {}+)", (effective_warmup * 10) / avg_steps_per_epoch ); warn!( " • Reduce --warmup-steps to {} (10% of total)", total_steps_estimate / 10 ); warn!(" • Set --warmup-steps 0 to disable warmup entirely"); } info!( " • Warmup steps: {}K ({:.1}% of training, Rainbow DQN random exploration)", effective_warmup / 1000, warmup_ratio * 100.0 ); } else { info!(" • Warmup steps: 0 (disabled for short training runs)"); } // Validate initial capital if opts.initial_capital < 1000.0 { eprintln!("āŒ Error: initial_capital must be >= $1,000 (got: ${:.2})", opts.initial_capital); eprintln!(" Use --initial-capital to specify a valid amount"); std::process::exit(1); } // Validate cash reserve percent if !(0.0..=100.0).contains(&opts.cash_reserve_percent) { eprintln!( "āŒ Error: cash_reserve_percent must be between 0.0 and 100.0 (got: {})", opts.cash_reserve_percent ); eprintln!(" Use --cash-reserve-percent to specify a valid amount"); std::process::exit(1); } // Setup graceful shutdown handler for containerized environments (RunPod, Docker, K8s) let shutdown_flag = Arc::new(AtomicBool::new(false)); let shutdown_clone = shutdown_flag.clone(); tokio::spawn(async move { let ctrl_c = signal::ctrl_c(); #[cfg(unix)] { use tokio::signal::unix::{signal, SignalKind}; let mut sigterm = signal(SignalKind::terminate()).expect("Failed to setup SIGTERM handler"); tokio::select! { _ = ctrl_c => { info!("šŸ›‘ Received Ctrl+C, initiating graceful shutdown..."); } _ = sigterm.recv() => { info!("šŸ›‘ Received SIGTERM, initiating graceful shutdown..."); } } } #[cfg(not(unix))] { ctrl_c.await.expect("Failed to listen for Ctrl+C"); info!("šŸ›‘ Received Ctrl+C, initiating graceful shutdown..."); } shutdown_clone.store(true, Ordering::Relaxed); }); info!("āœ… Graceful shutdown handler registered (Ctrl+C / SIGTERM)"); // Determine early stopping (enabled by default, unless --no-early-stopping is specified) let early_stopping_enabled = !opts.no_early_stopping; info!( " • Early stopping: {}", if early_stopping_enabled { "enabled" } else { "disabled" } ); if early_stopping_enabled { info!(" - Q-value floor: {}", opts.q_value_floor); info!(" - Min loss improvement: {}%", opts.min_loss_improvement); info!(" - Plateau window: {} epochs", opts.plateau_window); info!( " - Min epochs before stopping: {}", opts.min_epochs_before_stopping ); } // Create output and checkpoint directories let output_path = PathBuf::from(&opts.output_dir); let checkpoint_path = if let Some(ref dir) = opts.checkpoint_dir { PathBuf::from(dir) } else { output_path.clone() }; if !output_path.exists() { std::fs::create_dir_all(&output_path).context("Failed to create output directory")?; info!("āœ… Created output directory: {}", opts.output_dir); } if !checkpoint_path.exists() && checkpoint_path != output_path { std::fs::create_dir_all(&checkpoint_path) .context("Failed to create checkpoint directory")?; info!( "āœ… Created checkpoint directory: {}", checkpoint_path.display() ); } if opts.checkpoint_dir.is_some() { info!(" • Checkpoint directory: {}", checkpoint_path.display()); } // Configure DQN hyperparameters with optimal values from hyperopt (42 trials, 2025-10-31) let hyperparams = DQNHyperparameters { learning_rate: opts.learning_rate, batch_size: opts.batch_size, gamma: opts.gamma, epsilon_start: opts.epsilon_start, epsilon_end: opts.epsilon_end, epsilon_decay: opts.epsilon_decay, buffer_size: opts.buffer_size, min_replay_size: opts.min_replay_size, // Configurable min replay size epochs: opts.epochs, checkpoint_frequency: opts.checkpoint_frequency, early_stopping_enabled, q_value_floor: opts.q_value_floor, min_loss_improvement_pct: opts.min_loss_improvement, plateau_window: opts.plateau_window, min_epochs_before_stopping: opts.min_epochs_before_stopping, // NOW CONFIGURABLE! hold_penalty: -0.001, // Bug #3 fix: Enable gradient clipping to prevent Q-value collapse gradient_clip_norm: Some(10.0), // Conservative clipping at max_norm=10.0 // Bug #3 fix: Enable Huber loss for robustness to outliers use_huber_loss: true, huber_delta: 1.0, // Enable Double DQN to reduce overestimation bias use_double_dqn: true, // HOLD penalty weight (Bug #3 fix) hold_penalty_weight: opts.hold_penalty_weight, // Configurable via CLI // Price movement threshold for HOLD penalty (2% price movement) movement_threshold: 0.02, // Wave 14 Agent 32: Preprocessing configuration enable_preprocessing: !opts.no_preprocessing, // Enabled by default, disable with --no-preprocessing preprocessing_window: opts.preprocessing_window, preprocessing_clip_sigma: opts.preprocessing_clip_sigma, // Target update configuration (reverted to hard updates for stability) tau: opts.tau, // CLI-configurable (default: 1.0 = hard updates) target_update_mode: if opts.soft_updates { TargetUpdateMode::Soft } else { TargetUpdateMode::Hard }, // P2-B Enhancement: Cash reserve requirement cash_reserve_percent: opts.cash_reserve_percent, // Configurable via CLI target_update_frequency: 10000, // Hard update frequency (every 10K steps) // Rainbow DQN warmup warmup_steps: effective_warmup, // Adaptive warmup (0 for <200K, scaled 200K-1M, 80K for >1M) // P2-A Enhancement initial_capital: opts.initial_capital, // WAVE 16S: Adaptive Risk Management (ALL ENABLED BY DEFAULT) enable_kelly_sizing: true, enable_volatility_epsilon: true, enable_risk_adjusted_rewards: true, kelly_fractional: 0.5, kelly_max_fraction: 0.25, kelly_min_trades: 20, volatility_window: 20, // WAVE 35: Advanced Features (ALL ENABLED BY DEFAULT) enable_regime_qnetwork: true, enable_compliance: true, // WAVE 16: Core Risk Management (ALL ENABLED BY DEFAULT) enable_drawdown_monitoring: true, enable_position_limits: true, enable_circuit_breaker: true, // Wave 16 Portfolio Features (ALL ENABLED BY DEFAULT) enable_action_masking: true, enable_entropy_regularization: true, enable_stress_testing: true, }; // Configure alternative bar sampling (Wave B) let bar_sampling = match opts.bar_method.as_str() { "tick" => BarSamplingMethod::TickBars(opts.bar_threshold.unwrap_or(100.0) as usize), "volume" => BarSamplingMethod::VolumeBars(opts.bar_threshold.unwrap_or(10000.0)), "dollar" => BarSamplingMethod::DollarBars(opts.bar_threshold.unwrap_or(2_000_000.0)), "imbalance" => BarSamplingMethod::ImbalanceBars(opts.bar_threshold.unwrap_or(1000.0)), "run" => BarSamplingMethod::RunBars(opts.bar_threshold.unwrap_or(50.0) as usize), _ => BarSamplingMethod::TimeBars, }; info!("āœ… Bar sampling configured: {:?}", bar_sampling); // Create DQN trainer let mut trainer = DQNTrainer::new(hyperparams).context("Failed to create DQN trainer")?; // Note: DQN trainer will need to accept bar_sampling parameter // This requires updating DQNTrainer to use DbnSequenceLoader info!("āœ… DQN trainer initialized"); // Setup checkpoint manager let checkpoint_config = CheckpointConfig { base_dir: output_path.clone(), max_checkpoints_per_model: 10, auto_cleanup: true, validate_checksums: true, ..Default::default() }; let _checkpoint_manager = CheckpointManager::new(checkpoint_config).context("Failed to create checkpoint manager")?; info!("āœ… Checkpoint manager initialized (max 10 checkpoints, auto-cleanup enabled)"); // Create checkpoint callback with interruption handling let checkpoint_dir_for_callback = opts .checkpoint_dir .clone() .unwrap_or_else(|| opts.output_dir.clone()); let shutdown_check = shutdown_flag.clone(); let checkpoint_callback = move |epoch: usize, model_data: Vec, is_best: bool| -> Result { // Check if shutdown was requested let interrupted = shutdown_check.load(Ordering::Relaxed); let filename = if is_best { // Best model checkpoint (overwrites previous best) "dqn_best_model.safetensors".to_string() } else if interrupted { format!("dqn_interrupted_epoch{}.safetensors", epoch) } else { // Periodic checkpoint format!("dqn_epoch_{}.safetensors", epoch) }; let checkpoint_path = PathBuf::from(&checkpoint_dir_for_callback).join(filename); // Save checkpoint to disk std::fs::write(&checkpoint_path, &model_data) .context(format!("Failed to save checkpoint: {:?}", checkpoint_path))?; let checkpoint_type = if is_best { "šŸŽ‰ BEST" } else if interrupted { "āš ļø INTERRUPTED" } else { "šŸ’¾ PERIODIC" }; info!( "{} Checkpoint saved: {} ({} bytes)", checkpoint_type, checkpoint_path.display(), model_data.len() ); Ok(checkpoint_path.to_string_lossy().to_string()) }; // Train the model info!("\nšŸ‹ļø Starting training...\n"); let start_time = std::time::Instant::now(); let metrics = if let Some(ref parquet_path) = opts.parquet_file { info!("Using Parquet file: {}", parquet_path); trainer .train_from_parquet(parquet_path, checkpoint_callback) .await .context("Training from Parquet failed")? } else { info!("Using DBN directory: {}", opts.data_dir); trainer .train(&opts.data_dir, checkpoint_callback) .await .context("Training failed")? }; let training_duration = start_time.elapsed(); // Check if training was interrupted if shutdown_flag.load(Ordering::Relaxed) { info!("\nāš ļø Training was interrupted by shutdown signal"); info!("šŸ’¾ Interrupted checkpoint saved, safe to terminate"); info!("šŸ“Š Partial training metrics:"); info!(" • Epochs completed: {}", metrics.epochs_trained); info!( " • Training time: {:.1}s ({:.1} min)", metrics.training_time_seconds, metrics.training_time_seconds / 60.0 ); return Ok(()); } // Print final metrics info!("\nāœ… Training completed successfully!"); info!("\nšŸ“Š Final Metrics:"); info!(" • Final loss: {:.6}", metrics.loss); info!(" • Epochs trained: {}", metrics.epochs_trained); info!( " • Training time: {:.1}s ({:.1} min)", metrics.training_time_seconds, metrics.training_time_seconds / 60.0 ); info!( " • Actual elapsed time: {:.1}s (includes data loading + overhead)", training_duration.as_secs_f64() ); info!( " • Convergence: {}", if metrics.convergence_achieved { "āœ… Yes" } else { "āŒ No" } ); // Additional metrics from training if let Some(avg_q_value) = metrics.additional_metrics.get("avg_q_value") { info!(" • Average Q-value: {:.4}", avg_q_value); } if let Some(final_epsilon) = metrics.additional_metrics.get("final_epsilon") { info!(" • Final epsilon: {:.4}", final_epsilon); } if let Some(grad_norm) = metrics.additional_metrics.get("avg_gradient_norm") { info!(" • Average gradient norm: {:.6}", grad_norm); } // Save final model let final_model_path = output_path.join(format!("dqn_final_epoch{}.safetensors", opts.epochs)); info!("\nšŸ’¾ Saving final model to: {}", final_model_path.display()); // Get final model state let final_checkpoint_data = trainer .serialize_model() .await .context("Failed to serialize final model")?; std::fs::write(&final_model_path, &final_checkpoint_data) .context("Failed to save final model")?; info!( "āœ… Final model saved: {} ({} bytes)", final_model_path.display(), final_checkpoint_data.len() ); info!("\nšŸŽ‰ DQN training complete!"); info!("šŸ“ Model files saved to: {}", opts.output_dir); Ok(()) }