//! 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 clap::Parser; use std::path::PathBuf; use tracing::{info, warn}; use tracing_subscriber::FmtSubscriber; use ml::data_loaders::BarSamplingMethod; use ml::features::extraction::{extract_ml_features, OHLCVBar}; use ml::real_data_loader::RealDataLoader; use ml::trainers::ppo::{PpoHyperparameters, PpoTrainer, PpoTrainingMetrics}; #[derive(Debug, Parser)] #[command(name = "train_ppo", about = "Train PPO model on real market data")] struct Opts { /// Number of training epochs (default: 20 for policy convergence) #[arg(long, default_value = "20")] epochs: usize, /// Learning rate #[arg(long, default_value = "0.0003")] learning_rate: f64, /// Batch size (max 230 for RTX 3050 Ti 4GB) #[arg(long, default_value = "64")] batch_size: 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")] data_dir: String, /// Symbol to train on (ZN.FUT has ~29K bars) #[arg(long, default_value = "ZN.FUT")] symbol: String, /// 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, /// Minimum value loss improvement percentage for plateau detection #[arg(long, default_value = "2.0")] min_value_loss_improvement: f64, /// Minimum explained variance threshold #[arg(long, default_value = "0.4")] min_explained_variance: f64, /// Plateau detection window size (epochs) #[arg(long, default_value = "30")] plateau_window: usize, /// Alternative bar sampling method (time, tick, volume, dollar, imbalance, run) #[arg(long, default_value = "time")] bar_method: String, /// Bar sampling threshold (tick count, volume, dollar value, imbalance, or run length) #[arg(long)] bar_threshold: Option, } #[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")?; 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); info!(" • Bar sampling method: {}", opts.bar_method); if let Some(threshold) = opts.bar_threshold { info!(" • Bar threshold: {}", threshold); } // 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); } // 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); // Load real market data from DBN files info!("\nšŸ“Š Loading real market data from DBN files..."); let mut loader = RealDataLoader::new(&opts.data_dir); // Note: RealDataLoader will need to accept bar_sampling parameter // This requires updating RealDataLoader to use alternative bar sampling 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 using 225-feature extraction pipeline (Wave C) // Features 0-4: OHLCV (normalized) // Features 5-14: Technical indicators (10) // Features 15-74: Price patterns (60) // Features 75-114: Volume patterns (40) // Features 115-164: Microstructure proxies (50) // Features 165-174: Time-based features (10) // Features 175-200: Statistical features (26) // Features 201-224: Wave D regime detection (24) info!("\nšŸ—ļø Extracting 225-dimensional feature vectors..."); // Convert RealDataLoader bars to OHLCVBar format for feature extraction let ohlcv_bars: Vec = bars .iter() .map(|bar| OHLCVBar { timestamp: bar.timestamp, open: bar.open, high: bar.high, low: bar.low, close: bar.close, volume: bar.volume, }) .collect(); // Extract 225-dimensional feature vectors (requires 50-bar warmup) let feature_vectors = extract_ml_features(&ohlcv_bars) .context("Failed to extract 225-dimensional features")?; info!( "āœ… Extracted {} feature vectors (dim=225, warmup bars skipped=50)", feature_vectors.len() ); // Convert FeatureVector ([f64; 225]) to Vec> for PPO trainer let state_dim = 225; // Updated from 16 to 225 let market_data: Vec> = feature_vectors .iter() .map(|fv| fv.iter().map(|&v| v as f32).collect()) .collect(); // 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(()) }