//! Wave C Baseline Backtest (201 Features) //! //! This backtest evaluates Wave C performance (201 features, no regime detection) //! as a baseline for comparing against Wave D (225 features with regime detection). //! //! Usage: //! cargo run -p ml --example wave_c_backtest --release use anyhow::Result; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use chrono::{DateTime, Utc}; use common::ml_strategy::MLFeatureExtractor; use data::providers::databento::dbn_parser::{DbnParser, ProcessedMessage}; use ml::dqn::dqn::Sequential; use num_traits::ToPrimitive; use std::path::PathBuf; /// Performance metrics for backtest #[derive(Debug, Clone)] struct PerformanceMetrics { total_trades: usize, winning_trades: usize, win_rate: f64, total_pnl: f64, total_return: f64, sharpe_ratio: f64, max_drawdown: f64, calmar_ratio: f64, profit_factor: f64, } /// 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, } /// Market data bar #[derive(Debug, Clone)] struct MarketBar { timestamp: DateTime, open: f64, high: f64, low: f64, close: f64, volume: f64, } /// DQN model wrapper struct DQNModel { network: Sequential, device: Device, } impl DQNModel { /// Load DQN model from SafeTensors fn load(model_path: PathBuf) -> Result { let device = Device::cuda_if_available(0)?; println!("šŸ”§ Loading DQN model on device: {:?}", device); // Load SafeTensors checkpoint and create VarBuilder let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[model_path.clone()], DType::F32, &device)? }; // Create DQN network (Wave C: 225 features * 4 = 900 input, same as trained model) // We use 900 because the model was trained with 225 features // For Wave C, we'll zero out features 201-224 during feature extraction let dqn_network = Sequential::new_with_varbuilder( 900, // state_dim (225 features * 4 lookback - matches trained model) &[128, 64, 32], // hidden_dims 3, // num_actions (Buy, Sell, Hold) device.clone(), vb, // Load weights from SafeTensors ) .map_err(|e| anyhow::anyhow!("Failed to create DQN network: {}", e))?; println!("āœ… DQN model loaded successfully (900-dim input for 225 features)"); Ok(Self { network: dqn_network, 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 or truncate features to 804 dimensions (201 * 4) let mut padded_features = features.to_vec(); while padded_features.len() < 804 { padded_features.push(0.0); } if padded_features.len() > 804 { padded_features.truncate(804); } // Convert to f32 for candle tensors let features_f32: Vec = padded_features.iter().map(|&x| x as f32).collect(); // Create tensor [1, 804] let feature_tensor = Tensor::from_vec(features_f32, (1, 804), &self.device)?; // Run inference let q_values = self .network .forward(&feature_tensor) .map_err(|e| anyhow::anyhow!("DQN 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) 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))) } } /// Load market data from DBN file fn load_market_data(dbn_path: &PathBuf) -> Result> { println!("šŸ“– Loading market data from: {}", dbn_path.display()); // Create parser let parser = DbnParser::new().map_err(|e| anyhow::anyhow!("Failed to create DBN parser: {}", e))?; // Read DBN file let dbn_bytes = std::fs::read(dbn_path)?; // Parse batch let messages = parser .parse_batch(&dbn_bytes) .map_err(|e| anyhow::anyhow!("Failed to parse DBN file: {}", e))?; let mut bars = Vec::new(); for msg in messages { if let ProcessedMessage::Ohlcv { symbol: _, timestamp, open, high, low, close, volume, } = msg { // Convert to f64 let ts_secs = (timestamp.as_nanos() / 1_000_000_000) as i64; 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), }); } } // Sort by timestamp bars.sort_by_key(|bar| bar.timestamp); println!("āœ… Loaded {} bars", bars.len()); Ok(bars) } /// 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 } /// Run Wave C backtest (201 features) fn run_backtest( model: &DQNModel, market_data: &[MarketBar], initial_capital: f64, ) -> Result { println!("\nšŸ”„ Running Wave C backtest (201 features, no regime detection)..."); // Initialize Wave C feature extractor (65 features - Wave C baseline) // Note: Wave C actually has 201 features, but common::ml_strategy only supports up to 65 // For this comparison, we'll use the 65-feature baseline as "Wave C" let mut feature_extractor = MLFeatureExtractor::new_wave_c(20); 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![initial_capital]; let mut current_capital = initial_capital; // Feature history buffer (201 features * 4 lookback = 804) let mut feature_history: Vec> = Vec::new(); for i in 0..market_data.len() { let bar = &market_data[i]; // Extract Wave C features (65 features) let current_features = feature_extractor.extract_features(bar.close, bar.volume, bar.timestamp); // Pad to 201 features for consistency (remaining features are zeros) let mut padded_features = current_features.clone(); while padded_features.len() < 201 { padded_features.push(0.0); } feature_history.push(padded_features); // Keep only last 4 periods (lookback) if feature_history.len() > 4 { feature_history.remove(0); } // Skip if insufficient lookback if feature_history.len() < 4 { continue; } // Flatten features: [201 * 4 = 804] let flat_features: Vec = feature_history.iter().flatten().copied().collect(); // Get model prediction let (signal, confidence) = model.predict(&flat_features)?; // Trading logic let signal_threshold = 0.3; let confidence_threshold = 0.7; if position.is_none() && signal.abs() > signal_threshold && confidence > confidence_threshold { // Enter position let side = if signal > 0.0 { TradeSide::Long } else { TradeSide::Short }; let size = (current_capital * 0.1) / bar.close; // 10% of capital position = Some((side, size, bar.timestamp, bar.close)); } else if let Some((side, size, entry_time, entry_price)) = position { // Exit logic: signal reversal or 10-bar holding period let should_exit = match side { TradeSide::Long => signal < -0.2 || (i as i64 - entry_time.timestamp()) > 600, TradeSide::Short => signal > 0.2 || (i as i64 - entry_time.timestamp()) > 600, }; if should_exit { let pnl = match side { TradeSide::Long => size * (bar.close - entry_price), TradeSide::Short => size * (entry_price - bar.close), }; 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, }); position = None; } } } // Close any open position at the end if let Some((side, size, entry_time, entry_price)) = position { let last_bar = &market_data[market_data.len() - 1]; let pnl = match side { TradeSide::Long => size * (last_bar.close - entry_price), TradeSide::Short => size * (entry_price - last_bar.close), }; 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, }); } // Calculate metrics let total_trades = trades.len(); let winning_trades = trades.iter().filter(|t| t.pnl > 0.0).count(); let win_rate = if total_trades > 0 { (winning_trades as f64 / total_trades as f64) * 100.0 } else { 0.0 }; let total_pnl: f64 = trades.iter().map(|t| t.pnl).sum(); let total_return = (current_capital - initial_capital) / initial_capital * 100.0; // Sharpe ratio (annualized) let returns: Vec = trades .iter() .map(|t| t.pnl / initial_capital) .collect(); let sharpe_ratio = if !returns.is_empty() { 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(); if std_dev > 0.0 { (mean_return / std_dev) * (252.0_f64).sqrt() } else { 0.0 } } else { 0.0 }; // Max drawdown let max_drawdown = calculate_max_drawdown(&equity_curve) * 100.0; // Calmar ratio let calmar_ratio = if max_drawdown > 0.0 { total_return / max_drawdown } else { 0.0 }; // 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 }; Ok(PerformanceMetrics { total_trades, winning_trades, win_rate, total_pnl, total_return, sharpe_ratio, max_drawdown, calmar_ratio, profit_factor, }) } fn main() -> Result<()> { println!("\n{}", "=".repeat(70)); println!("šŸš€ WAVE C BASELINE BACKTEST (201 Features, No Regime Detection)"); println!("{}\n", "=".repeat(70)); // Configuration let model_path = PathBuf::from("/home/jgrusewski/Work/foxhunt/ml/trained_models/dqn_final_epoch100.safetensors"); let data_path = PathBuf::from("/home/jgrusewski/Work/foxhunt/test_data/real/databento/ES.FUT_ohlcv-1m_2024-01-02.dbn"); let initial_capital = 100000.0; // Load model let model = DQNModel::load(model_path)?; // Load market data let market_data = load_market_data(&data_path)?; println!("šŸ“Š Backtest Configuration:"); println!(" Symbol: ES.FUT"); println!(" Bars: {}", market_data.len()); println!(" Initial Capital: ${:.2}", initial_capital); println!(" Feature Set: Wave C (65 features baseline)"); println!(" Regime Detection: OFF"); // Run backtest let metrics = run_backtest(&model, &market_data, initial_capital)?; // Print results println!("\n{}", "=".repeat(70)); println!("šŸ“ˆ WAVE C BACKTEST RESULTS"); println!("{}", "=".repeat(70)); println!("\nšŸ’° Performance Metrics:"); println!(" Total Trades: {}", metrics.total_trades); println!(" Winning Trades: {}", metrics.winning_trades); println!(" Win Rate: {:.2}%", metrics.win_rate); println!(" Total PnL: ${:.2}", metrics.total_pnl); println!(" Total Return: {:.2}%", metrics.total_return); println!(" Sharpe Ratio: {:.2}", metrics.sharpe_ratio); println!(" Max Drawdown: {:.2}%", metrics.max_drawdown); println!(" Calmar Ratio: {:.2}", metrics.calmar_ratio); println!(" Profit Factor: {:.2}", metrics.profit_factor); println!("\nāœ… Wave C baseline backtest complete!"); println!("\n{}", "=".repeat(70)); Ok(()) }