//! TFT (Temporal Fusion Transformer) Training Example //! //! Trains a TFT model for time series forecasting and saves checkpoints. //! //! # Usage //! //! ```bash //! # Train with default parameters (100 epochs) //! cargo run -p ml --example train_tft --release --features cuda //! //! # Custom configuration //! cargo run -p ml --example train_tft --release --features cuda -- \ //! --epochs 500 \ //! --batch-size 32 \ //! --hidden-dim 256 //! ``` use anyhow::{Context, Result}; use ndarray::{Array1, Array2, Array3}; use std::path::PathBuf; use std::sync::Arc; use structopt::StructOpt; use tokio::sync::mpsc; use tracing::info; use tracing_subscriber::FmtSubscriber; use ml::checkpoint::FileSystemStorage; use ml::trainers::tft::{TFTTrainer, TFTTrainerConfig}; use ml::tft::training::TFTDataLoader; #[derive(Debug, StructOpt)] #[structopt(name = "train_tft", about = "Train TFT model on time series data")] struct Opts { /// Number of training epochs #[structopt(long, default_value = "100")] epochs: usize, /// Learning rate #[structopt(long, default_value = "0.001")] learning_rate: f64, /// Batch size (max 32 for 4GB VRAM) #[structopt(long, default_value = "32")] batch_size: usize, /// Hidden dimension #[structopt(long, default_value = "256")] hidden_dim: usize, /// Number of attention heads #[structopt(long, default_value = "8")] num_attention_heads: usize, /// Lookback window #[structopt(long, default_value = "60")] lookback_window: usize, /// Forecast horizon #[structopt(long, default_value = "10")] forecast_horizon: usize, /// Output directory for trained model #[structopt(long, default_value = "ml/trained_models")] output_dir: String, /// Use GPU #[structopt(long)] use_gpu: bool, /// Verbose logging #[structopt(short, long)] verbose: bool, } #[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 TFT Training"); info!("Configuration:"); info!(" • Epochs: {}", opts.epochs); info!(" • Learning rate: {}", opts.learning_rate); info!(" • Batch size: {}", opts.batch_size); info!(" • Hidden dimension: {}", opts.hidden_dim); info!(" • Attention heads: {}", opts.num_attention_heads); info!(" • Lookback window: {}", opts.lookback_window); info!(" • Forecast horizon: {}", opts.forecast_horizon); info!(" • GPU enabled: {}", opts.use_gpu); info!(" • Output directory: {}", opts.output_dir); // 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 TFT trainer let trainer_config = TFTTrainerConfig { epochs: opts.epochs, learning_rate: opts.learning_rate, batch_size: opts.batch_size, hidden_dim: opts.hidden_dim, num_attention_heads: opts.num_attention_heads, dropout_rate: 0.1, lstm_layers: 2, quantiles: vec![0.1, 0.5, 0.9], lookback_window: opts.lookback_window, forecast_horizon: opts.forecast_horizon, use_gpu: opts.use_gpu, checkpoint_dir: opts.output_dir.clone(), }; // Create checkpoint storage let storage = std::sync::Arc::new(FileSystemStorage::new(output_path.clone())); // Create TFT trainer let mut trainer = TFTTrainer::new(trainer_config.clone(), storage) .context("Failed to create TFT trainer")?; info!("āœ… TFT trainer initialized"); // Generate synthetic time series data info!("\nšŸ“Š Generating training data..."); let num_train_samples = 3200; // 100 batches of size 32 let num_val_samples = 320; // 10 batches of size 32 let train_loader = generate_data_loader( num_train_samples, opts.batch_size, opts.lookback_window, opts.forecast_horizon, true, // shuffle training data )?; let val_loader = generate_data_loader( num_val_samples, opts.batch_size, opts.lookback_window, opts.forecast_horizon, false, // don't shuffle validation data )?; info!("āœ… Generated {} training samples, {} validation samples", num_train_samples, num_val_samples); // Setup progress callback let (progress_tx, mut progress_rx) = mpsc::unbounded_channel(); trainer.set_progress_callback(progress_tx); // Spawn progress monitor task let monitor_task = tokio::spawn(async move { while let Some(progress) = progress_rx.recv().await { if progress.current_epoch % 10 == 0 { info!("{}", progress.message); if let Some(loss) = progress.metrics.get("train_loss") { info!(" • Train loss: {:.6}", loss); } if let Some(val_loss) = progress.metrics.get("val_loss") { info!(" • Val loss: {:.6}", val_loss); } if let Some(rmse) = progress.metrics.get("rmse") { info!(" • RMSE: {:.6}", rmse); } } } }); // Train the model info!("\nšŸ‹ļø Starting training...\n"); let start_time = std::time::Instant::now(); let final_metrics = trainer .train(train_loader, val_loader) .await .context("Training failed")?; let training_duration = start_time.elapsed(); // Wait for progress monitor to finish drop(trainer); // Drop trainer to close progress channel let _ = monitor_task.await; // Print final metrics info!("\nāœ… Training completed successfully!"); info!("\nšŸ“Š Final Metrics:"); info!(" • Training loss: {:.6}", final_metrics.train_loss); info!(" • Validation loss: {:.6}", final_metrics.val_loss); info!(" • Quantile loss: {:.6}", final_metrics.quantile_loss); info!(" • RMSE: {:.6}", final_metrics.rmse); info!(" • Attention entropy: {:.4}", final_metrics.attention_entropy); info!(" • Training time: {:.1}s ({:.1} min)", final_metrics.training_time_seconds, final_metrics.training_time_seconds / 60.0); info!("\nšŸ’¾ Model checkpoints saved to: {}", opts.output_dir); info!("\nšŸŽ‰ TFT training complete!"); Ok(()) } /// Generate synthetic data loader for training/validation fn generate_data_loader( num_samples: usize, batch_size: usize, lookback_window: usize, forecast_horizon: usize, shuffle: bool, ) -> Result { use ndarray::Array1; // Generate synthetic data samples let mut data = Vec::with_capacity(num_samples); for i in 0..num_samples { // Static features: [num_static_features] = [10] let static_features = Array1::from_shape_fn(10, |j| { (i as f64 * 0.1 + j as f64 * 0.01) }); // Historical features: [lookback_window, num_hist_features] = [60, 50] let historical_features = Array2::from_shape_fn( (lookback_window, 50), |(t, f)| { (i as f64 * 0.1 + t as f64 * 0.01 + f as f64 * 0.001).sin() } ); // Future features: [forecast_horizon, num_fut_features] = [10, 10] let future_features = Array2::from_shape_fn( (forecast_horizon, 10), |(t, f)| { (i as f64 * 0.1 + (lookback_window + t) as f64 * 0.01 + f as f64 * 0.001).cos() } ); // Targets: [forecast_horizon] = [10] let targets = Array1::from_shape_fn( forecast_horizon, |t| { (i as f64 * 0.1 + (lookback_window + t) as f64 * 0.01).sin() * 100.0 } ); data.push((static_features, historical_features, future_features, targets)); } Ok(TFTDataLoader::new(data, batch_size, shuffle)) }