//! Comprehensive backtesting for all trained ML models //! //! This example loads all available trained models and runs backtesting with real market data. //! It generates performance metrics including Sharpe ratio, win rate, max drawdown, and PnL. //! //! Usage: //! cargo run -p ml --example comprehensive_model_backtest --release use anyhow::Result; use chrono::{DateTime, Utc}; use num_traits::ToPrimitive; use serde::{Deserialize, Serialize}; use std::path::PathBuf; use candle_core::{Device, Tensor, DType}; use candle_nn::VarBuilder; use data::providers::databento::dbn_parser::{DbnParser, ProcessedMessage}; use ml::dqn::dqn::Sequential; use ml::ppo::ppo::PolicyNetwork; /// Backtesting configuration #[derive(Debug, Clone)] struct BacktestConfig { /// Model checkpoint path model_path: PathBuf, /// Data directory data_dir: PathBuf, /// Symbol to test symbol: String, /// Start date start_date: DateTime, /// End date end_date: DateTime, /// Initial capital initial_capital: f64, /// Position size position_size: f64, } /// Performance metrics for a backtest #[derive(Debug, Clone, Serialize, Deserialize)] struct PerformanceMetrics { /// Model name model_name: String, /// Model type (DQN or PPO) model_type: String, /// Epoch number epoch: u32, /// Total trades total_trades: usize, /// Winning trades winning_trades: usize, /// Win rate percentage win_rate: f64, /// Total PnL total_pnl: f64, /// Sharpe ratio sharpe_ratio: f64, /// Max drawdown percentage max_drawdown: f64, /// Calmar ratio (return / max drawdown) calmar_ratio: f64, /// Average trade duration (minutes) avg_trade_duration: f64, /// Profit factor (gross profit / gross loss) profit_factor: f64, /// Trade frequency (trades per 1000 bars) trade_frequency: f64, /// Start date start_date: String, /// End date end_date: String, } /// Trade record #[derive(Debug, Clone)] struct Trade { entry_time: DateTime, exit_time: DateTime, entry_price: f64, exit_price: f64, side: TradeSide, pnl: f64, size: f64, } #[derive(Debug, Clone, Copy)] enum TradeSide { Long, Short, } /// Model type enum enum ModelType { DQN(Sequential), PPO(PolicyNetwork), } /// Simple model inference wrapper struct ModelInference { model_name: String, model_type: ModelType, device: Device, } impl ModelInference { /// Load DQN model from SafeTensors fn load_dqn(model_name: String, model_path: PathBuf) -> Result { let device = Device::cuda_if_available(0)?; println!("šŸ”§ Loading DQN model: {} on device: {:?}", model_name, device); // Load SafeTensors checkpoint let _vb = unsafe { VarBuilder::from_mmaped_safetensors(&[model_path.clone()], DType::F32, &device)? }; // Create DQN network architecture (64 -> 128 -> 64 -> 32 -> 3) let dqn_network = Sequential::new( 64, // state_dim (16 features * 4 = 64) &[128, 64, 32], // hidden_dims 3, // num_actions (Buy, Sell, Hold) device.clone(), ).map_err(|e| anyhow::anyhow!("Failed to create DQN network: {}", e))?; println!("āœ… DQN model loaded successfully"); Ok(Self { model_name, model_type: ModelType::DQN(dqn_network), device, }) } /// Load PPO model from SafeTensors fn load_ppo(model_name: String, model_path: PathBuf) -> Result { let device = Device::cuda_if_available(0)?; println!("šŸ”§ Loading PPO model: {} on device: {:?}", model_name, device); // Load SafeTensors checkpoint let _vb = unsafe { VarBuilder::from_mmaped_safetensors(&[model_path.clone()], DType::F32, &device)? }; // Create PPO actor network (64 -> 128 -> 64 -> 3) let ppo_actor = PolicyNetwork::new( 64, // state_dim &[128, 64], // hidden_dims 3, // num_actions device.clone(), ).map_err(|e| anyhow::anyhow!("Failed to create PPO network: {}", e))?; println!("āœ… PPO model loaded successfully"); Ok(Self { model_name, model_type: ModelType::PPO(ppo_actor), device, }) } /// Predict trading signal from features /// Returns: (signal_strength: -1.0 to 1.0, confidence: 0.0 to 1.0) fn predict(&self, features: &[f64]) -> Result<(f64, f64)> { // Pad features to 64 dimensions if needed let mut padded_features = features.to_vec(); while padded_features.len() < 64 { padded_features.push(0.0); } if padded_features.len() > 64 { padded_features.truncate(64); } // Convert to f32 for candle tensors let features_f32: Vec = padded_features.iter().map(|&x| x as f32).collect(); // Create tensor [1, 64] let feature_tensor = Tensor::from_vec(features_f32, (1, 64), &self.device)?; // Run inference based on model type let q_values = match &self.model_type { ModelType::DQN(network) => { network.forward(&feature_tensor) .map_err(|e| anyhow::anyhow!("DQN forward pass failed: {}", e))? } ModelType::PPO(actor) => { actor.forward(&feature_tensor) .map_err(|e| anyhow::anyhow!("PPO forward pass failed: {}", e))? } }; // Get action probabilities let q_vec = q_values.to_vec2::()?; let actions = &q_vec[0]; // [Buy, Sell, Hold] // Convert action values to signal (-1 to 1) // Buy = 1.0, Sell = -1.0, Hold = 0.0 let buy_strength = actions[0] as f64; let sell_strength = actions[1] as f64; let hold_strength = actions[2] as f64; // Normalize to -1 to 1 range let signal = if buy_strength > sell_strength && buy_strength > hold_strength { (buy_strength - hold_strength).min(1.0) } else if sell_strength > buy_strength && sell_strength > hold_strength { -(sell_strength - hold_strength).min(1.0) } else { 0.0 }; // Confidence based on action strength difference let max_action = buy_strength.max(sell_strength).max(hold_strength); let confidence = (max_action - hold_strength).abs().min(1.0); Ok((signal, confidence.max(0.5))) } } /// Feature extractor for market data struct FeatureExtractor { price_history: Vec, volume_history: Vec, lookback: usize, } impl FeatureExtractor { fn new(lookback: usize) -> Self { Self { price_history: Vec::with_capacity(lookback), volume_history: Vec::with_capacity(lookback), lookback, } } fn extract_features(&mut self, price: f64, volume: f64) -> Vec { self.price_history.push(price); self.volume_history.push(volume); // Keep only lookback period if self.price_history.len() > self.lookback { self.price_history.remove(0); self.volume_history.remove(0); } let mut features = Vec::new(); if self.price_history.len() < 2 { return vec![0.0; 10]; // Return zeros if insufficient data } let current_price = price; let prev_price = self.price_history[self.price_history.len() - 2]; // 1. Price momentum (% change) let price_change = (current_price - prev_price) / prev_price; features.push(price_change); // 2. SMA ratio (price vs 10-period SMA) if self.price_history.len() >= 10 { let sma: f64 = self.price_history.iter().rev().take(10).sum::() / 10.0; let sma_ratio = (current_price - sma) / sma; features.push(sma_ratio); } else { features.push(0.0); } // 3. RSI (14-period) let rsi = self.calculate_rsi(14); features.push(rsi); // 4. Volume ratio if self.volume_history.len() >= 2 { let curr_vol = volume; let prev_vol = self.volume_history[self.volume_history.len() - 2]; let vol_ratio = if prev_vol > 0.0 { (curr_vol - prev_vol) / prev_vol } else { 0.0 }; features.push(vol_ratio); } else { features.push(0.0); } // 5. Volatility (20-period std dev of returns) if self.price_history.len() >= 20 { let returns: Vec = self.price_history .windows(2) .map(|w| (w[1] - w[0]) / w[0]) .collect(); let mean = returns.iter().sum::() / returns.len() as f64; let variance = returns.iter() .map(|r| (r - mean).powi(2)) .sum::() / returns.len() as f64; let volatility = variance.sqrt(); features.push(volatility); } else { features.push(0.0); } // Pad to 10 features while features.len() < 10 { features.push(0.0); } features } fn calculate_rsi(&self, period: usize) -> f64 { if self.price_history.len() < period + 1 { return 50.0; // Neutral RSI } let recent_prices: Vec = self.price_history .iter() .rev() .take(period + 1) .copied() .collect(); let mut gains = 0.0; let mut losses = 0.0; for i in 1..recent_prices.len() { let change = recent_prices[i-1] - recent_prices[i]; if change > 0.0 { gains += change; } else { losses += change.abs(); } } let avg_gain = gains / period as f64; let avg_loss = losses / period as f64; if avg_loss == 0.0 { return 100.0; } let rs = avg_gain / avg_loss; let rsi = 100.0 - (100.0 / (1.0 + rs)); rsi } } /// Run backtest for a model fn run_backtest(config: BacktestConfig, is_dqn: bool, epoch: u32, total_bars: usize) -> Result { println!("\n{}", "=".repeat(60)); println!("šŸŽÆ Starting backtest: {}", config.symbol); println!(" Model: {}", config.model_path.display()); println!(" Period: {} to {}", config.start_date, config.end_date); println!("{}\n", "=".repeat(60)); // Initialize model let model_name = config.model_path .file_stem() .and_then(|s| s.to_str()) .unwrap_or("unknown") .to_string(); let model = if is_dqn { ModelInference::load_dqn(model_name.clone(), config.model_path.clone())? } else { ModelInference::load_ppo(model_name.clone(), config.model_path.clone())? }; // Load market data println!("šŸ“Š Loading market data from: {}", config.data_dir.display()); let market_data = load_market_data(&config.data_dir, &config.symbol)?; if market_data.is_empty() { anyhow::bail!("No market data found for symbol: {}", config.symbol); } println!("āœ… Loaded {} bars", market_data.len()); // Initialize feature extractor let mut feature_extractor = FeatureExtractor::new(50); // Run backtest let mut trades = Vec::new(); let mut position: Option<(TradeSide, f64, DateTime, f64)> = None; // (side, size, entry_time, entry_price) let mut equity_curve = vec![config.initial_capital]; let mut current_capital = config.initial_capital; println!("šŸ”„ Running backtest simulation..."); for (i, bar) in market_data.iter().enumerate() { // Extract features let features = feature_extractor.extract_features(bar.close, bar.volume); // Get model prediction let (signal, confidence) = model.predict(&features)?; // Only trade if confidence is high enough if confidence < 0.6 { continue; } // Check for entry signal if position.is_none() { if signal > 0.5 { // Enter long position = Some((TradeSide::Long, config.position_size, bar.timestamp, bar.close)); if i % 100 == 0 { println!(" šŸ“ˆ LONG entry at {:.2} (signal: {:.3}, confidence: {:.3})", bar.close, signal, confidence); } } else if signal < -0.5 { // Enter short position = Some((TradeSide::Short, config.position_size, bar.timestamp, bar.close)); if i % 100 == 0 { println!(" šŸ“‰ SHORT entry at {:.2} (signal: {:.3}, confidence: {:.3})", bar.close, signal, confidence); } } } else if let Some((side, size, entry_time, entry_price)) = position { // Check for exit signal let should_exit = match side { TradeSide::Long => signal < -0.3, // Exit long on negative signal TradeSide::Short => signal > 0.3, // Exit short on positive signal }; if should_exit { // Calculate PnL let pnl = match side { TradeSide::Long => (bar.close - entry_price) * size, TradeSide::Short => (entry_price - bar.close) * size, }; current_capital += pnl; equity_curve.push(current_capital); trades.push(Trade { entry_time, exit_time: bar.timestamp, entry_price, exit_price: bar.close, side, pnl, size, }); if i % 100 == 0 { println!(" āœ… Exit at {:.2}, PnL: {:.2} (signal: {:.3})", bar.close, pnl, signal); } position = None; } } if i % 500 == 0 && i > 0 { let progress = (i as f64 / market_data.len() as f64) * 100.0; println!(" Progress: {:.1}% ({} trades)", progress, trades.len()); } } // Close any open position at the end if let Some((side, size, entry_time, entry_price)) = position { let last_bar = market_data.last().unwrap(); let pnl = match side { TradeSide::Long => (last_bar.close - entry_price) * size, TradeSide::Short => (entry_price - last_bar.close) * size, }; current_capital += pnl; equity_curve.push(current_capital); trades.push(Trade { entry_time, exit_time: last_bar.timestamp, entry_price, exit_price: last_bar.close, side, pnl, size, }); } println!("\nāœ… Backtest complete! {} trades executed", trades.len()); // Calculate performance metrics calculate_performance_metrics(model_name, trades, equity_curve, config, is_dqn, epoch, total_bars) } /// Calculate performance metrics from trades fn calculate_performance_metrics( model_name: String, trades: Vec, equity_curve: Vec, config: BacktestConfig, is_dqn: bool, epoch: u32, total_bars: usize, ) -> Result { let model_type = if is_dqn { "DQN" } else { "PPO" }; if trades.is_empty() { return Ok(PerformanceMetrics { model_name, model_type: model_type.to_string(), epoch, total_trades: 0, winning_trades: 0, win_rate: 0.0, total_pnl: 0.0, sharpe_ratio: 0.0, max_drawdown: 0.0, calmar_ratio: 0.0, avg_trade_duration: 0.0, profit_factor: 0.0, trade_frequency: 0.0, start_date: config.start_date.to_rfc3339(), end_date: config.end_date.to_rfc3339(), }); } // Basic metrics let total_trades = trades.len(); let winning_trades = trades.iter().filter(|t| t.pnl > 0.0).count(); let win_rate = (winning_trades as f64 / total_trades as f64) * 100.0; let total_pnl: f64 = trades.iter().map(|t| t.pnl).sum(); // Trade duration let avg_trade_duration: f64 = trades.iter() .map(|t| (t.exit_time - t.entry_time).num_minutes() as f64) .sum::() / total_trades as f64; // Profit factor let gross_profit: f64 = trades.iter().filter(|t| t.pnl > 0.0).map(|t| t.pnl).sum(); let gross_loss: f64 = trades.iter().filter(|t| t.pnl < 0.0).map(|t| t.pnl.abs()).sum(); let profit_factor = if gross_loss > 0.0 { gross_profit / gross_loss } else { if gross_profit > 0.0 { f64::INFINITY } else { 0.0 } }; // Sharpe ratio (annualized) let returns: Vec = trades.iter().map(|t| t.pnl / config.initial_capital).collect(); let mean_return = returns.iter().sum::() / returns.len() as f64; let variance = returns.iter() .map(|r| (r - mean_return).powi(2)) .sum::() / returns.len() as f64; let std_dev = variance.sqrt(); // Annualize (assume 252 trading days) let sharpe_ratio = if std_dev > 0.0 { (mean_return / std_dev) * (252.0_f64).sqrt() } else { 0.0 }; // Max drawdown let max_drawdown = calculate_max_drawdown(&equity_curve); // Calmar ratio let total_return = (equity_curve.last().unwrap() - config.initial_capital) / config.initial_capital; let calmar_ratio = if max_drawdown > 0.0 { total_return / max_drawdown } else { 0.0 }; // Trade frequency (trades per 1000 bars) let trade_frequency = if total_bars > 0 { (total_trades as f64 / total_bars as f64) * 1000.0 } else { 0.0 }; Ok(PerformanceMetrics { model_name, model_type: model_type.to_string(), epoch, total_trades, winning_trades, win_rate, total_pnl, sharpe_ratio, max_drawdown: max_drawdown * 100.0, // Convert to percentage calmar_ratio, avg_trade_duration, profit_factor, trade_frequency, start_date: config.start_date.to_rfc3339(), end_date: config.end_date.to_rfc3339(), }) } /// Calculate maximum drawdown from equity curve fn calculate_max_drawdown(equity_curve: &[f64]) -> f64 { let mut max_drawdown = 0.0; let mut peak = equity_curve[0]; for &equity in equity_curve { if equity > peak { peak = equity; } let drawdown = (peak - equity) / peak; if drawdown > max_drawdown { max_drawdown = drawdown; } } max_drawdown } /// Market data bar #[derive(Debug, Clone)] struct MarketBar { timestamp: DateTime, open: f64, high: f64, low: f64, close: f64, volume: f64, } /// Load market data from DBN files fn load_market_data(data_dir: &PathBuf, symbol: &str) -> Result> { println!("šŸ” Searching for {} data in {:?}", symbol, data_dir); // Find DBN files for the symbol let dbn_files: Vec = std::fs::read_dir(data_dir)? .filter_map(|entry| entry.ok()) .map(|entry| entry.path()) .filter(|path| { path.extension().and_then(|s| s.to_str()) == Some("dbn") && path.file_name() .and_then(|s| s.to_str()) .map(|s| s.contains(symbol)) .unwrap_or(false) }) .collect(); if dbn_files.is_empty() { anyhow::bail!("No DBN files found for symbol {} in {:?}", symbol, data_dir); } println!("šŸ“ Found {} DBN files for {}", dbn_files.len(), symbol); // Load actual DBN data let mut all_bars = Vec::new(); // Create parser let parser = DbnParser::new() .map_err(|e| anyhow::anyhow!("Failed to create DBN parser: {}", e))?; for dbn_file in dbn_files.iter().take(4) { println!("šŸ“– Reading: {}", dbn_file.display()); // Read DBN file let dbn_bytes = std::fs::read(dbn_file)?; // Parse batch let messages = parser.parse_batch(&dbn_bytes) .map_err(|e| anyhow::anyhow!("Failed to parse DBN file: {}", e))?; let mut file_bars = Vec::new(); for msg in messages { if let ProcessedMessage::Ohlcv { symbol: _, timestamp, open, high, low, close, volume } = msg { // Convert Price and Decimal to f64 let ts_secs = (timestamp.as_nanos() / 1_000_000_000) as i64; file_bars.push(MarketBar { timestamp: DateTime::from_timestamp(ts_secs, 0) .unwrap_or_else(|| Utc::now()), open: open.to_f64(), high: high.to_f64(), low: low.to_f64(), close: close.to_f64(), volume: volume.to_f64().unwrap_or(0.0), }); } } println!(" Loaded {} bars from {}", file_bars.len(), dbn_file.file_name().unwrap().to_str().unwrap()); all_bars.extend(file_bars); } // Sort by timestamp all_bars.sort_by_key(|bar| bar.timestamp); println!("āœ… Total bars loaded: {}", all_bars.len()); Ok(all_bars) } fn main() -> Result<()> { println!("\n{}", "=".repeat(70)); println!("šŸš€ COMPREHENSIVE ML MODEL BACKTESTING - ALL 101 CHECKPOINTS"); println!("{}\n", "=".repeat(70)); // Get project root let project_root = std::env::current_dir()?; let data_dir = project_root.join("test_data/real/databento/ml_training_small"); let model_dir = project_root.join("ml/trained_models/production"); let results_dir = project_root.join("results"); // Create results directory std::fs::create_dir_all(&results_dir)?; // Symbol to test let symbol = "6E.FUT"; // Pre-load market data once (shared across all models) println!("šŸ“Š Pre-loading market data from: {}", data_dir.display()); let market_data = load_market_data(&data_dir, symbol)?; let total_bars = market_data.len(); println!("āœ… Loaded {} bars for testing\n", total_bars); // Run backtests for all checkpoints let mut all_results = Vec::new(); // Test DQN checkpoints (epochs 10-500, every 10 epochs = 50 checkpoints) println!("\n{}", "=".repeat(70)); println!("šŸ”µ TESTING DQN CHECKPOINTS (50 models)"); println!("{}\n", "=".repeat(70)); let dqn_dir = model_dir.join("dqn_real_data"); for epoch in (10..=500).step_by(10) { let model_path = dqn_dir.join(format!("dqn_epoch_{}.safetensors", epoch)); if !model_path.exists() { println!("āš ļø DQN epoch {} not found: {}", epoch, model_path.display()); continue; } println!("Testing DQN epoch {}... ({}/50)", epoch, epoch / 10); let config = BacktestConfig { model_path: model_path.clone(), data_dir: data_dir.clone(), symbol: symbol.to_string(), start_date: chrono::Utc::now() - chrono::Duration::days(90), end_date: chrono::Utc::now(), initial_capital: 100_000.0, position_size: 1.0, }; match run_backtest(config, true, epoch, total_bars) { Ok(metrics) => { println!(" āœ… DQN epoch {}: {} trades, Sharpe {:.3}, Win rate {:.1}%", epoch, metrics.total_trades, metrics.sharpe_ratio, metrics.win_rate); all_results.push(metrics); } Err(e) => { println!(" āŒ DQN epoch {} failed: {}", epoch, e); } } } // Test PPO checkpoints (epochs 10-500, every 10 epochs = 50 checkpoints) println!("\n{}", "=".repeat(70)); println!("🟢 TESTING PPO CHECKPOINTS (50 models)"); println!("{}\n", "=".repeat(70)); let ppo_dir = model_dir.join("ppo_real_data"); for epoch in (10..=500).step_by(10) { let model_path = ppo_dir.join(format!("ppo_actor_epoch_{}.safetensors", epoch)); if !model_path.exists() { println!("āš ļø PPO epoch {} not found: {}", epoch, model_path.display()); continue; } println!("Testing PPO epoch {}... ({}/50)", epoch, epoch / 10); let config = BacktestConfig { model_path: model_path.clone(), data_dir: data_dir.clone(), symbol: symbol.to_string(), start_date: chrono::Utc::now() - chrono::Duration::days(90), end_date: chrono::Utc::now(), initial_capital: 100_000.0, position_size: 1.0, }; match run_backtest(config, false, epoch, total_bars) { Ok(metrics) => { println!(" āœ… PPO epoch {}: {} trades, Sharpe {:.3}, Win rate {:.1}%", epoch, metrics.total_trades, metrics.sharpe_ratio, metrics.win_rate); all_results.push(metrics); } Err(e) => { println!(" āŒ PPO epoch {} failed: {}", epoch, e); } } } // Save results to JSON let timestamp = chrono::Utc::now().format("%Y%m%d_%H%M%S"); let results_file = results_dir.join(format!("comprehensive_backtest_results_{}.json", timestamp)); let json = serde_json::to_string_pretty(&all_results)?; std::fs::write(&results_file, json)?; println!("\n{}", "=".repeat(70)); println!("āœ… Backtesting complete! Tested {} models", all_results.len()); println!("šŸ“Š Results saved to: {}", results_file.display()); println!("{}\n", "=".repeat(70)); // Print comprehensive summary print_comprehensive_summary(&all_results); // Save summary CSV save_summary_csv(&all_results, &results_dir)?; Ok(()) } fn print_comprehensive_summary(results: &[PerformanceMetrics]) { println!("{}", "=".repeat(90)); println!("šŸ“Š COMPREHENSIVE SUMMARY - ALL 101 MODELS"); println!("{}\n", "=".repeat(90)); if results.is_empty() { println!("āš ļø No results to display"); return; } // Separate DQN and PPO results let dqn_results: Vec<_> = results.iter().filter(|m| m.model_type == "DQN").collect(); let ppo_results: Vec<_> = results.iter().filter(|m| m.model_type == "PPO").collect(); // Print DQN summary println!("šŸ”µ DQN MODELS ({} total)", dqn_results.len()); println!("{}", "-".repeat(90)); println!("{:<12} {:>8} {:>10} {:>10} {:>12} {:>12} {:>12}", "Epoch", "Trades", "Win Rate", "Sharpe", "PnL", "Drawdown", "Trade Freq"); println!("{}", "-".repeat(90)); let mut dqn_sorted = dqn_results.clone(); dqn_sorted.sort_by(|a, b| b.sharpe_ratio.partial_cmp(&a.sharpe_ratio).unwrap_or(std::cmp::Ordering::Equal)); for metrics in dqn_sorted.iter().take(10) { println!( "{:<12} {:>8} {:>9.1}% {:>10.3} ${:>10.2} {:>11.2}% {:>12.1}", format!("Epoch {}", metrics.epoch), metrics.total_trades, metrics.win_rate, metrics.sharpe_ratio, metrics.total_pnl, metrics.max_drawdown, metrics.trade_frequency ); } println!("\n"); // Print PPO summary println!("🟢 PPO MODELS ({} total)", ppo_results.len()); println!("{}", "-".repeat(90)); println!("{:<12} {:>8} {:>10} {:>10} {:>12} {:>12} {:>12}", "Epoch", "Trades", "Win Rate", "Sharpe", "PnL", "Drawdown", "Trade Freq"); println!("{}", "-".repeat(90)); let mut ppo_sorted = ppo_results.clone(); ppo_sorted.sort_by(|a, b| b.sharpe_ratio.partial_cmp(&a.sharpe_ratio).unwrap_or(std::cmp::Ordering::Equal)); for metrics in ppo_sorted.iter().take(10) { println!( "{:<12} {:>8} {:>9.1}% {:>10.3} ${:>10.2} {:>11.2}% {:>12.1}", format!("Epoch {}", metrics.epoch), metrics.total_trades, metrics.win_rate, metrics.sharpe_ratio, metrics.total_pnl, metrics.max_drawdown, metrics.trade_frequency ); } println!("\n"); // Overall best models println!("{}", "=".repeat(90)); println!("šŸ† TOP 5 MODELS (All Types - Ranked by Sharpe Ratio)"); println!("{}", "=".repeat(90)); let mut all_sorted = results.to_vec(); all_sorted.sort_by(|a, b| b.sharpe_ratio.partial_cmp(&a.sharpe_ratio).unwrap_or(std::cmp::Ordering::Equal)); println!("{:<15} {:>8} {:>8} {:>10} {:>10} {:>12} {:>12}", "Model", "Epoch", "Trades", "Win Rate", "Sharpe", "PnL", "Trade Freq"); println!("{}", "-".repeat(90)); for (rank, metrics) in all_sorted.iter().take(5).enumerate() { println!( "{}. {:<12} {:>8} {:>8} {:>9.1}% {:>10.3} ${:>10.2} {:>12.1}", rank + 1, metrics.model_type, metrics.epoch, metrics.total_trades, metrics.win_rate, metrics.sharpe_ratio, metrics.total_pnl, metrics.trade_frequency ); } println!("\n"); // Statistics println!("{}", "=".repeat(90)); println!("šŸ“ˆ STATISTICAL SUMMARY"); println!("{}", "=".repeat(90)); let avg_sharpe: f64 = results.iter().map(|m| m.sharpe_ratio).sum::() / results.len() as f64; let avg_win_rate: f64 = results.iter().map(|m| m.win_rate).sum::() / results.len() as f64; let avg_trades: f64 = results.iter().map(|m| m.total_trades as f64).sum::() / results.len() as f64; let best_sharpe = results.iter().max_by(|a, b| { a.sharpe_ratio.partial_cmp(&b.sharpe_ratio).unwrap_or(std::cmp::Ordering::Equal) }).unwrap(); let best_win_rate = results.iter().max_by(|a, b| { a.win_rate.partial_cmp(&b.win_rate).unwrap_or(std::cmp::Ordering::Equal) }).unwrap(); let best_pnl = results.iter().max_by(|a, b| { a.total_pnl.partial_cmp(&b.total_pnl).unwrap_or(std::cmp::Ordering::Equal) }).unwrap(); println!("Average Sharpe Ratio: {:.3}", avg_sharpe); println!("Average Win Rate: {:.1}%", avg_win_rate); println!("Average Trades: {:.1}", avg_trades); println!(); println!("Best Sharpe: {:.3} ({} Epoch {})", best_sharpe.sharpe_ratio, best_sharpe.model_type, best_sharpe.epoch); println!("Best Win Rate: {:.1}% ({} Epoch {})", best_win_rate.win_rate, best_win_rate.model_type, best_win_rate.epoch); println!("Best PnL: ${:.2} ({} Epoch {})", best_pnl.total_pnl, best_pnl.model_type, best_pnl.epoch); println!("\n"); } fn save_summary_csv(results: &[PerformanceMetrics], results_dir: &PathBuf) -> Result<()> { let timestamp = chrono::Utc::now().format("%Y%m%d_%H%M%S"); let csv_file = results_dir.join(format!("backtest_summary_{}.csv", timestamp)); let mut csv_content = String::from( "model_type,epoch,total_trades,winning_trades,win_rate,total_pnl,sharpe_ratio,max_drawdown,calmar_ratio,avg_trade_duration,profit_factor,trade_frequency\n" ); for metrics in results { csv_content.push_str(&format!( "{},{},{},{},{:.2},{:.2},{:.4},{:.2},{:.4},{:.2},{:.4},{:.2}\n", metrics.model_type, metrics.epoch, metrics.total_trades, metrics.winning_trades, metrics.win_rate, metrics.total_pnl, metrics.sharpe_ratio, metrics.max_drawdown, metrics.calmar_ratio, metrics.avg_trade_duration, metrics.profit_factor, metrics.trade_frequency )); } std::fs::write(&csv_file, csv_content)?; println!("šŸ“Š CSV summary saved to: {}", csv_file.display()); Ok(()) }