//! TFT Hyperparameter Optimization Validation //! //! This example validates TFT hyperopt with a small dataset to identify //! any bugs similar to those found in MAMBA-2: //! 1. LR schedule bugs (dynamic parameter calculation) //! 2. Device transfer issues (CPU vs GPU tensor mismatches) //! 3. Tensor rank errors (unconditional squeeze operations) //! 4. Accuracy/metric calculation errors //! 5. Cache-related bugs //! //! ## Usage //! //! ```bash //! cargo run -p ml --example validate_tft_hyperopt --release --features cuda //! ``` use anyhow::Result; use ml::hyperopt::adapters::tft::{TFTParams, TFTTrainer}; use ml::hyperopt::{ArgminOptimizer, HyperparameterOptimizable}; use tracing::{info, Level}; fn main() -> Result<()> { // Initialize tracing tracing_subscriber::fmt() .with_max_level(Level::INFO) .with_target(false) .init(); info!("========================================"); info!("TFT Hyperopt Validation"); info!("========================================"); info!("Goal: Identify MAMBA-2-style bugs"); info!("Dataset: ES_FUT_small.parquet (~200 samples)"); info!("Trials: 3 × 5 epochs (~30 seconds)"); info!(""); // Create trainer with small dataset and few epochs info!("Creating TFT trainer..."); let parquet_file = "test_data/ES_FUT_small.parquet"; let epochs = 5; // Small epoch count for fast validation let trainer = TFTTrainer::new(parquet_file, epochs)?; info!("TFT Configuration:"); info!(" Input features: 54 (Wave D)"); info!(" Sequence length: 60"); info!(" Prediction horizon: 10"); info!(" Quantiles: 3 (0.1, 0.5, 0.9)"); info!(" Epochs per trial: {}", epochs); info!(""); // Create optimizer with minimal trials info!("Initializing optimizer..."); let optimizer = ArgminOptimizer::builder() .max_trials(3) // Just 3 trials for validation .n_initial(2) // 2 random + 1 optimized .seed(42) .build(); // Run optimization info!(""); info!("Starting optimization..."); info!("Expected runtime: ~30 seconds"); info!(""); let start = std::time::Instant::now(); let result = optimizer.optimize(trainer)?; let elapsed = start.elapsed(); // Display results info!(""); info!("========================================"); info!("Validation Complete!"); info!("========================================"); info!(""); info!("Runtime: {:.1}s", elapsed.as_secs_f64()); info!(""); info!("Best Hyperparameters:"); info!(" Learning rate: {:.6}", result.best_params.learning_rate); info!(" Batch size: {}", result.best_params.batch_size); info!(" Hidden size: {}", result.best_params.hidden_size); info!(" Attention heads: {}", result.best_params.num_heads); info!(" Dropout: {:.3}", result.best_params.dropout); info!(""); info!("Performance:"); info!(" Best validation loss: {:.6}", result.best_objective); info!(" Total trials: {}", result.all_trials.len()); info!(""); // Analyze trial results info!("Trial Analysis:"); let mut success_count = 0; for (i, trial) in result.all_trials.iter().enumerate() { let status = if trial.objective < 10.0 { success_count += 1; "✓ SUCCESS" } else { "✗ FAILED" }; info!( " Trial {}: Loss={:.6} LR={:.6} BS={} Hidden={} Heads={} {}", i + 1, trial.objective, trial.params.learning_rate, trial.params.batch_size, trial.params.hidden_size, trial.params.num_heads, status ); } let success_rate = (success_count as f64 / result.all_trials.len() as f64) * 100.0; info!(""); info!( "Success Rate: {}/{} ({:.1}%)", success_count, result.all_trials.len(), success_rate ); if success_rate == 100.0 { info!("✓ VALIDATION PASSED: All trials successful (like PPO)"); } else { info!("⚠ VALIDATION ISSUE: Some trials failed (investigate bugs)"); } info!(""); info!("========================================"); info!("Bug Check (vs MAMBA-2 issues):"); info!("========================================"); // Check for specific bug patterns info!("1. LR Schedule: Check if total_decay_steps is calculated dynamically"); info!(" → Review ml/src/hyperopt/adapters/tft.rs"); info!(""); info!("2. Device Transfer: Check validation functions use .to_device()"); info!(" → Review ml/src/trainers/tft.rs forward() calls"); info!(""); info!("3. Tensor Rank: Check for unconditional .squeeze() operations"); info!(" → Review ml/src/tft/mod.rs tensor operations"); info!(""); info!("4. Metric Calculation: Check loss computation for edge cases"); info!(" → Review quantile loss calculation in TFT trainer"); info!(""); info!("5. Cache: Check attention cache is cleared properly"); info!(" → Review TFT cache management in trainer"); info!(""); if success_rate == 100.0 && result.best_objective < 1.0 { info!("✓ TFT Hyperopt: PRODUCTION READY (zero bugs found)"); } else if success_rate == 100.0 { info!("⚠ TFT Hyperopt: TRIALS PASS but loss high (tune parameters)"); } else { info!("✗ TFT Hyperopt: BUGS DETECTED (fix before deployment)"); } Ok(()) }