//! PPO Training Example with Real DataBento Market Data //! //! Trains a PPO model on real market data from DBN files with: //! - Real OHLCV data + technical indicators //! - Actual PnL-based rewards //! - GAE advantages on real price trajectories //! - Policy convergence validation (KL divergence > 0) //! //! # Usage //! //! ```bash //! # Train with default parameters (20 epochs) //! cargo run -p ml --example train_ppo --release --features cuda //! //! # Custom epochs and output path //! cargo run -p ml --example train_ppo --release --features cuda -- \ //! --epochs 50 \ //! --output-dir ml/trained_models \ //! --data-dir test_data/real/databento //! ``` use anyhow::{Context, Result}; use std::path::PathBuf; use structopt::StructOpt; use tracing::{info, warn}; use tracing_subscriber::FmtSubscriber; use ml::real_data_loader::RealDataLoader; use ml::trainers::ppo::{PpoHyperparameters, PpoTrainer, PpoTrainingMetrics}; #[derive(Debug, StructOpt)] #[structopt(name = "train_ppo", about = "Train PPO model on real market data")] struct Opts { /// Number of training epochs (default: 20 for policy convergence) #[structopt(long, default_value = "20")] epochs: usize, /// Learning rate #[structopt(long, default_value = "0.0003")] learning_rate: f64, /// Batch size (max 230 for RTX 3050 Ti 4GB) #[structopt(long, default_value = "64")] batch_size: usize, /// Output directory for trained model #[structopt(long, default_value = "ml/trained_models")] output_dir: String, /// Data directory containing DBN files #[structopt(long, default_value = "test_data/real/databento")] data_dir: String, /// Symbol to train on (ZN.FUT has ~29K bars) #[structopt(long, default_value = "ZN.FUT")] symbol: String, /// Verbose logging #[structopt(short, long)] verbose: bool, /// Enable early stopping (recommended, use --no-early-stopping to disable) #[structopt(long)] early_stopping: bool, /// Disable early stopping #[structopt(long)] no_early_stopping: bool, /// Minimum value loss improvement percentage for plateau detection #[structopt(long, default_value = "2.0")] min_value_loss_improvement: f64, /// Minimum explained variance threshold #[structopt(long, default_value = "0.4")] min_explained_variance: f64, /// Plateau detection window size (epochs) #[structopt(long, default_value = "30")] plateau_window: usize, } #[tokio::main] async fn main() -> Result<()> { // Parse CLI options let opts = Opts::from_args(); // 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")?; info!("šŸš€ Starting PPO Training with Real DataBento Data"); info!("Configuration:"); info!(" • Epochs: {}", opts.epochs); info!(" • Learning rate: {}", opts.learning_rate); info!(" • Batch size: {}", opts.batch_size); info!(" • GPU: CUDA MANDATORY (no CPU fallback)"); info!(" • Output directory: {}", opts.output_dir); info!(" • Data directory: {}", opts.data_dir); info!(" • Symbol: {}", opts.symbol); // 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!(" - Min value loss improvement: {}%", opts.min_value_loss_improvement); info!(" - Min explained variance: {}", opts.min_explained_variance); info!(" - Plateau window: {} epochs", opts.plateau_window); } // Create output directory let output_path = PathBuf::from(&opts.output_dir); if !output_path.exists() { std::fs::create_dir_all(&output_path) .context("Failed to create output directory")?; info!("āœ… Created output directory: {}", opts.output_dir); } // Load real market data from DBN files info!("\nšŸ“Š Loading real market data from DBN files..."); let mut loader = RealDataLoader::new(&opts.data_dir); let bars = loader.load_symbol_data(&opts.symbol).await .context(format!("Failed to load data for symbol: {}", opts.symbol))?; info!("āœ… Loaded {} OHLCV bars for {}", bars.len(), opts.symbol); // Extract features and indicators info!("\nšŸ”§ Extracting features and technical indicators..."); let features = loader.extract_features(&bars) .context("Failed to extract features")?; let indicators = loader.calculate_indicators(&bars) .context("Failed to calculate indicators")?; info!("āœ… Feature extraction complete:"); info!(" • OHLCV bars: {}", features.prices.len()); info!(" • Returns: {}", features.returns.len()); info!(" • Volume: {}", features.volume.len()); info!(" • Indicators: 10 technical indicators"); // Build PPO state vectors (OHLCV + indicators + returns) // State: [open, high, low, close, volume, rsi, macd, macd_signal, bb_upper, bb_middle, // bb_lower, atr, ema_fast, ema_slow, volume_ma, log_return] info!("\nšŸ—ļø Building PPO state vectors..."); let state_dim = 16; // 5 (OHLCV) + 10 (indicators) + 1 (return) let mut market_data = Vec::with_capacity(bars.len()); for i in 0..bars.len() { let mut state = Vec::with_capacity(state_dim); // OHLCV (normalized 0-1) state.extend_from_slice(&features.prices[i]); // Technical indicators (10 values) state.push(indicators.rsi[i]); state.push(indicators.macd[i]); state.push(indicators.macd_signal[i]); state.push(indicators.bb_upper[i]); state.push(indicators.bb_middle[i]); state.push(indicators.bb_lower[i]); state.push(indicators.atr[i]); state.push(indicators.ema_fast[i]); state.push(indicators.ema_slow[i]); state.push(indicators.volume_ma[i]); // Log return state.push(features.returns[i]); market_data.push(state); } info!("āœ… Built {} state vectors (dim={})", market_data.len(), state_dim); // Validate state dimensions if let Some(first_state) = market_data.first() { if first_state.len() != state_dim { return Err(anyhow::anyhow!( "State dimension mismatch: expected {}, got {}", state_dim, first_state.len() )); } } // Configure PPO hyperparameters let hyperparams = PpoHyperparameters { learning_rate: opts.learning_rate, batch_size: opts.batch_size, gamma: 0.99, clip_epsilon: 0.2, vf_coef: 0.5, ent_coef: 0.01, gae_lambda: 0.95, rollout_steps: 2048, minibatch_size: opts.batch_size, epochs: opts.epochs, early_stopping_enabled, min_value_loss_improvement_pct: opts.min_value_loss_improvement, min_explained_variance: opts.min_explained_variance, plateau_window: opts.plateau_window, min_epochs_before_stopping: 50, }; // Create PPO trainer with real data state dimension let trainer = PpoTrainer::new( hyperparams.clone(), state_dim, &opts.output_dir, true, // CUDA always required ).context("Failed to create PPO trainer")?; info!("āœ… PPO trainer initialized (state_dim={})", state_dim); // Create progress callback with convergence tracking let mut policy_updates = 0; let mut kl_divergence_history = Vec::new(); let progress_callback = |metrics: PpoTrainingMetrics| { // Track policy updates (KL divergence > 0 indicates policy changed) if metrics.kl_divergence > 0.0 { policy_updates += 1; } kl_divergence_history.push(metrics.kl_divergence); info!( "šŸ“Š Epoch {}/{}: policy_loss={:.4}, value_loss={:.4}, kl_div={:.6}, expl_var={:.4}, mean_reward={:.4}", metrics.epoch, hyperparams.epochs, metrics.policy_loss, metrics.value_loss, metrics.kl_divergence, metrics.explained_variance, metrics.mean_reward ); }; // Train the model info!("\nšŸ‹ļø Starting training...\n"); let start_time = std::time::Instant::now(); let final_metrics = trainer .train(market_data, progress_callback) .await .context("Training failed")?; let training_duration = start_time.elapsed(); // Print final metrics info!("\nāœ… Training completed successfully!"); info!("\nšŸ“Š Final Metrics:"); info!(" • Policy loss: {:.6}", final_metrics.policy_loss); info!(" • Value loss: {:.6}", final_metrics.value_loss); info!(" • KL divergence: {:.6}", final_metrics.kl_divergence); info!(" • Explained variance: {:.4}", final_metrics.explained_variance); info!(" • Mean reward: {:.4}", final_metrics.mean_reward); info!(" • Std reward: {:.4}", final_metrics.std_reward); info!(" • Entropy: {:.4}", final_metrics.entropy); info!(" • Training time: {:.1}s ({:.1} min)", training_duration.as_secs_f64(), training_duration.as_secs_f64() / 60.0); // Validate policy convergence info!("\nšŸ” Policy Convergence Analysis:"); info!(" • Total epochs: {}", hyperparams.epochs); info!(" • Policy updates (KL > 0): {}", policy_updates); info!(" • Policy update rate: {:.1}%", (policy_updates as f64 / hyperparams.epochs as f64) * 100.0); // Calculate KL divergence statistics let kl_mean = kl_divergence_history.iter().sum::() / kl_divergence_history.len() as f32; let kl_max = kl_divergence_history.iter().copied().fold(f32::NEG_INFINITY, f32::max); let kl_min = kl_divergence_history.iter().copied().fold(f32::INFINITY, f32::min); info!(" • KL divergence (mean): {:.6}", kl_mean); info!(" • KL divergence (max): {:.6}", kl_max); info!(" • KL divergence (min): {:.6}", kl_min); // Convergence validation if final_metrics.kl_divergence > 0.0 { info!(" āœ… PASS: Policy updates detected (KL divergence > 0)"); } else { warn!(" āš ļø WARN: No policy updates in final epoch (KL divergence = 0)"); warn!(" This may indicate learning rate too low or convergence"); } // Value function validation if final_metrics.explained_variance > 0.5 { info!(" āœ… PASS: Value network learning (explained variance > 0.5)"); } else { warn!(" āš ļø WARN: Value network may need tuning (explained variance < 0.5)"); } // Checkpoint is already saved by trainer (every 10 epochs) let final_checkpoint = output_path.join(format!("ppo_checkpoint_epoch_{}.safetensors", hyperparams.epochs)); info!("\nšŸ’¾ Final checkpoint saved to: {}", final_checkpoint.display()); info!("\nšŸŽ‰ PPO training complete with real DataBento data!"); info!("šŸ“ Model files saved to: {}", opts.output_dir); info!("\nšŸ“ˆ Training Summary:"); info!(" • Data source: Real DataBento OHLCV ({})", opts.symbol); info!(" • Training samples: {}", bars.len()); info!(" • State dimension: {}", state_dim); info!(" • Features: OHLCV + 10 technical indicators + log returns"); info!(" • Policy updates: {}/{} epochs ({:.1}%)", policy_updates, hyperparams.epochs, (policy_updates as f64 / hyperparams.epochs as f64) * 100.0); info!(" • Convergence: {}", if final_metrics.kl_divergence > 0.0 { "āœ… Achieved" } else { "āš ļø Check logs" }); Ok(()) }