Changes: - CLAUDE.md: Update OOM fix validation status - Add comprehensive documentation (30+ markdown reports) - LSTM encoder varmap bug fix (tft/lstm_encoder.rs:290) - Quantized LSTM layer matching fix (tft/quantized_lstm.rs) - Hyperopt paths module (ml/src/hyperopt/paths.rs) - Training path tests for all adapters (DQN, MAMBA-2, PPO, TFT) - Checkpoint integrity tests - Script cleanup: Remove 29 obsolete deployment scripts - Archive old scripts to scripts/archive/ - New deployment utilities: check_gpu_availability.py, monitor_hyperopt.sh Validation: - OOM fixes validated: 5/5 trials successful (pod b6kc3mc5lbjiro) - Batch-size-max 256 tested successfully - All hyperopt adapters working correctly 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude <noreply@anthropic.com>
159 lines
5.3 KiB
Rust
159 lines
5.3 KiB
Rust
//! 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: 225 (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(())
|
||
}
|