//! Simple DQN Trading Evaluation //! //! Minimal evaluation script that loads a trained DQN model and evaluates //! its trading performance on test data using the integrated evaluation engine. //! //! # Usage //! ```bash //! cargo run -p ml --example simple_dqn_eval --release --features cuda -- \ //! --model ml/trained_models/dqn_best_model.safetensors \ //! --data test_data/ES_FUT_180d.parquet //! ``` use anyhow::{Context, Result}; use candle_core::{Device, Tensor}; use clap::Parser; use ml::data_loaders::load_parquet_data; use ml::dqn::dqn::{WorkingDQN, WorkingDQNConfig}; use ml::evaluation::engine::{Action, EvaluationEngine}; use ml::evaluation::metrics::{OHLCVBar, PerformanceMetrics}; use ml::features::extraction::compute_dqn_features; use std::path::PathBuf; use tracing::{info, warn}; #[derive(Parser)] struct Args { /// Path to trained model checkpoint #[arg(long, default_value = "ml/trained_models/dqn_best_model.safetensors")] model: PathBuf, /// Path to test data (Parquet) #[arg(long, default_value = "test_data/ES_FUT_180d.parquet")] data: PathBuf, /// Initial capital #[arg(long, default_value_t = 100000.0)] capital: f32, /// Warmup bars to skip #[arg(long, default_value_t = 50)] warmup: usize, } fn main() -> Result<()> { // Initialize logging tracing_subscriber::fmt() .with_max_level(tracing::Level::INFO) .init(); let args = Args::parse(); info!("=== Trial #2 DQN Model Evaluation ==="); info!("Model: {}", args.model.display()); info!("Data: {}", args.data.display()); info!("Capital: ${}", args.capital); // Step 1: Load parquet data info!("Loading test data..."); let (features, bars) = load_parquet_data(&args.data, args.warmup) .with_context(|| format!("Failed to load {}", args.data.display()))?; info!("Loaded {} bars ({} features)", bars.len(), features.len()); // Validate data if features.is_empty() || bars.len() < args.warmup { anyhow::bail!( "Insufficient data: {} bars, {} warmup required", bars.len(), args.warmup ); } // Step 2: Create device let device = Device::cuda_if_available(0).context("Failed to create compute device")?; info!("Device: {:?}", device); // Step 3: Load DQN model info!("Loading DQN model from {}...", args.model.display()); // Create config matching training hyperparameters let config = WorkingDQNConfig { state_dim: 225, // Feature dimension action_dim: 3, // BUY, HOLD, SELL hidden_dim: 256, learning_rate: 0.000156, gamma: 0.97, epsilon: 0.01, // Greedy during evaluation epsilon_decay: 0.995, epsilon_min: 0.01, batch_size: 100, buffer_capacity: 642214, target_update_freq: 10, hold_penalty_weight: 1.0, movement_threshold: 0.02, gradient_clip_norm: 10.0, leaky_relu_alpha: 0.01, tau: 0.005, warmup_steps: 0, }; let mut dqn = WorkingDQN::new(config, device.clone()).context("Failed to create DQN model")?; // Load weights from checkpoint dqn.load(&args.model) .with_context(|| format!("Failed to load model from {}", args.model.display()))?; info!("✅ Model loaded successfully"); // Step 4: Run evaluation info!("Running backtest evaluation..."); let mut engine = EvaluationEngine::new(args.capital); for (idx, (feature, bar)) in features.iter().zip(bars.iter()).enumerate() { // Convert feature to DQN state (225-dim) let state = Tensor::from_slice(feature.as_slice(), (1, 225), &device) .context("Failed to create state tensor")?; // Get DQN action (greedy, no exploration during eval) let action_idx = dqn .select_action_greedy(&state) .context("Failed to select action")?; let action = Action::from(action_idx); // Process action through evaluation engine let ohlcv_bar = OHLCVBar { timestamp: bar.timestamp, open: bar.open, high: bar.high, low: bar.low, close: bar.close, volume: bar.volume, }; engine.process_bar(idx, &ohlcv_bar, action); // Log every 1000 bars if (idx + 1) % 1000 == 0 { info!("Processed {}/{} bars", idx + 1, features.len()); } } // Close any remaining position if let Some(last_bar) = bars.last() { let ohlcv_bar = OHLCVBar { timestamp: last_bar.timestamp, open: last_bar.open, high: last_bar.high, low: last_bar.low, close: last_bar.close, volume: last_bar.volume, }; engine.close_position(features.len() - 1, &ohlcv_bar); } // Step 5: Calculate metrics info!("Calculating performance metrics..."); let metrics = PerformanceMetrics::from_trades(&engine.trades, args.capital, &bars); let action_dist = engine.get_action_distribution(); // Step 6: Print report println!("\n{}", "=".repeat(70)); println!("TRIAL #2 DQN MODEL EVALUATION REPORT"); println!("{}", "=".repeat(70)); println!(); println!("MODEL DETAILS:"); println!(" Checkpoint: {}", args.model.display()); println!(" Training Epochs: 100 (best at epoch 61)"); println!(" Validation Loss: 8,017.93"); println!(); println!("HYPERPARAMETERS:"); println!(" Learning Rate: 0.000156"); println!(" Batch Size: 100"); println!(" Gamma: 0.97"); println!(" Buffer Size: 642,214"); println!(" Hold Penalty: 1.0"); println!(); println!("BACKTEST PERFORMANCE:"); println!(" Total Return: {:.2}%", metrics.total_return_pct); println!(" Sharpe Ratio: {:.2}", metrics.sharpe_ratio); println!(" Max Drawdown: {:.2}%", metrics.max_drawdown_pct); println!(" Win Rate: {:.2}%", metrics.win_rate * 100.0); println!(" Total Trades: {}", metrics.total_trades); println!(" Avg Trade P&L: ${:.2}", metrics.avg_trade_pnl); println!(" Final Equity: ${:.2}", metrics.final_equity); println!(" Total P&L: ${:.2}", metrics.total_pnl); println!(); println!("ACTION DISTRIBUTION:"); println!( " BUY: {:>6} ({:>5.2}%)", action_dist.buy_count, action_dist.buy_pct ); println!( " SELL: {:>6} ({:>5.2}%)", action_dist.sell_count, action_dist.sell_pct ); println!( " HOLD: {:>6} ({:>5.2}%)", action_dist.hold_count, action_dist.hold_pct ); println!(); println!("PRODUCTION READINESS:"); let sharpe_ok = metrics.sharpe_ratio >= 1.5; let win_rate_ok = metrics.win_rate >= 0.55; let drawdown_ok = metrics.max_drawdown_pct <= 20.0; let profitable = metrics.total_pnl > 0.0; println!( " {} Sharpe Ratio >= 1.5: {:.2}", if sharpe_ok { "✅" } else { "❌" }, metrics.sharpe_ratio ); println!( " {} Win Rate >= 55%: {:.2}%", if win_rate_ok { "✅" } else { "❌" }, metrics.win_rate * 100.0 ); println!( " {} Max Drawdown <= 20%: {:.2}%", if drawdown_ok { "✅" } else { "❌" }, metrics.max_drawdown_pct ); println!( " {} Profitable: ${:.2}", if profitable { "✅" } else { "❌" }, metrics.total_pnl ); println!(); let production_ready = sharpe_ok && win_rate_ok && drawdown_ok && profitable; if production_ready { println!("✅ PRODUCTION READY - All criteria met!"); } else { println!("⚠️ NOT PRODUCTION READY - Some criteria not met"); } println!("{}", "=".repeat(70)); println!(); Ok(()) }