//! Simplified DQN Real Training Validation //! //! Quick validation test using a small DBN file subset (5 files ~7,500 bars) use anyhow::Result; use ml::trainers::dqn::{DQNHyperparameters, DQNTrainer}; use std::fs; use std::path::Path; use tracing::{info, Level}; use tracing_subscriber::FmtSubscriber; #[tokio::main] async fn main() -> Result<()> { // Initialize logging let subscriber = FmtSubscriber::builder() .with_max_level(Level::INFO) .finish(); tracing::subscriber::set_global_default(subscriber)?; info!("========================================"); info!("DQN Real Training Quick Validation"); info!("========================================"); // Create temp directory with subset of files let temp_dir = "/tmp/dqn_test_data"; fs::create_dir_all(temp_dir)?; // Copy first 5 DBN files let source_dir = "test_data/real/databento/ml_training/"; let files: Vec<_> = fs::read_dir(source_dir)? .filter_map(|e| e.ok()) .filter(|e| e.path().extension().and_then(|s| s.to_str()) == Some("dbn")) .take(5) .collect(); info!("Copying {} DBN files to temp directory...", files.len()); for entry in files { let src = entry.path(); let dst = Path::new(temp_dir).join(entry.file_name()); fs::copy(&src, &dst)?; } // Configure for quick training (2 epochs, small batch) let mut hyperparams = DQNHyperparameters::conservative(); hyperparams.epochs = 2; hyperparams.batch_size = 32; hyperparams.buffer_size = 1_000; hyperparams.checkpoint_frequency = 1; hyperparams.early_stopping_enabled = false; info!("\nHyperparameters:"); info!(" Epochs: {}", hyperparams.epochs); info!(" Batch size: {}", hyperparams.batch_size); info!(" Learning rate: {}", hyperparams.learning_rate); // Create trainer let mut trainer = DQNTrainer::new(hyperparams)?; // Checkpoint callback (no-op) let checkpoint_callback = |epoch: usize, _: Vec| -> Result { Ok(format!("/tmp/dqn_test_epoch_{}.safetensors", epoch)) }; // Run training info!("\nStarting 2-epoch training...\n"); let start_time = std::time::Instant::now(); let metrics = trainer.train(temp_dir, checkpoint_callback).await?; let training_duration = start_time.elapsed(); // Analyze results info!("\n========================================"); info!("Training Complete - Results Analysis"); info!("========================================"); info!("\nFinal Metrics:"); info!(" Loss: {:.6}", metrics.loss); info!(" Epochs: {}", metrics.epochs_trained); info!(" Time: {:.2}s", training_duration.as_secs_f64()); let avg_q_value = metrics .additional_metrics .get("avg_q_value") .copied() .unwrap_or(0.0); let avg_grad_norm = metrics .additional_metrics .get("avg_gradient_norm") .copied() .unwrap_or(0.0); let final_epsilon = metrics .additional_metrics .get("final_epsilon") .copied() .unwrap_or(0.1); info!("\nDQN Metrics:"); info!(" Q-value: {:.4}", avg_q_value); info!(" Grad norm: {:.6}", avg_grad_norm); info!(" Epsilon: {:.4}", final_epsilon); // Validation info!("\n========================================"); info!("Validation Checks"); info!("========================================"); let mut passed = true; // Check 1: Loss is not hardcoded 0.5 if (metrics.loss - 0.5).abs() > 1e-6 { info!("✅ Loss is dynamic ({:.6})", metrics.loss); } else { info!("❌ Loss is hardcoded (0.5)"); passed = false; } // Check 2: Q-value is not hardcoded 10.0 if (avg_q_value - 10.0).abs() > 1e-6 { info!("✅ Q-value is dynamic ({:.4})", avg_q_value); } else { info!("❌ Q-value is hardcoded (10.0)"); passed = false; } // Check 3: Gradient norm is not hardcoded 0.01 if (avg_grad_norm - 0.01).abs() > 1e-6 { info!("✅ Gradient norm is dynamic ({:.6})", avg_grad_norm); } else { info!("❌ Gradient norm is hardcoded (0.01)"); passed = false; } // Check 4: Training completed if metrics.epochs_trained == 2 { info!("✅ Completed 2 epochs"); } else { info!("❌ Expected 2 epochs, got {}", metrics.epochs_trained); passed = false; } // Cleanup info!("\nCleaning up temp directory..."); fs::remove_dir_all(temp_dir)?; // Final verdict info!("\n========================================"); if passed { info!("✅ ALL VALIDATION CHECKS PASSED"); info!(" DQN uses REAL Q-learning algorithm!"); } else { info!("❌ VALIDATION FAILED"); } info!("========================================"); if !passed { anyhow::bail!("Validation failed"); } Ok(()) }