Files
foxhunt/ml/examples/train_ppo.rs
jgrusewski f946dcd952 feat: Wave 2 - Update MEDIUM RISK files (225→54 features)
WAVE 22: All examples, benchmarks, and data loaders updated

Files Modified (41 files):
- DQN examples: 7 files (train_dqn, evaluate_dqn, validate_dqn, etc.)
- PPO examples: 6 files (train_ppo, continuous_ppo, benchmark_ppo, etc.)
- TFT examples: 9 files (train_tft, validate_tft, benchmark_tft, etc.)
- MAMBA-2 examples: 3 files (train_mamba2, verify_dimensions, etc.)
- Benchmarks: 5 files (cuda_speedup, weight_caching, future_decoder, etc.)
- Data loaders: 7 files (parquet_utils, dbn_sequence_loader, tlob_loader, etc.)
- Integration: 4 files (load_parquet_data, streaming loaders, etc.)

Key Changes:
- state_dim: 225 → 54 (DQN, PPO)
- input_dim: 225 → 54 (TFT)
- d_model: 225 → 54 (MAMBA-2)
- Memory: 1.8KB → 0.43KB per vector (76% reduction)
- All tensor shapes updated: (batch, 225) → (batch, 54)

Agents Deployed: 5 parallel agents
Validation: cargo check PASSING

Generated with Claude Code

Co-Authored-By: Claude <noreply@anthropic.com>
2025-11-23 00:57:17 +01:00

407 lines
14 KiB
Rust
Raw Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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 mimalloc allocator for 10-25% performance improvement
#[cfg(feature = "mimalloc-allocator")]
use mimalloc::MiMalloc;
#[cfg(feature = "mimalloc-allocator")]
#[global_allocator]
static GLOBAL: MiMalloc = MiMalloc;
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 (512 recommended for value network stability, prevents -23.56 explained variance failure)
#[arg(long, default_value = "512")]
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<f64>,
}
#[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")?;
#[cfg(feature = "mimalloc-allocator")]
info!("🚀 Using mimalloc allocator for improved performance");
#[cfg(not(feature = "mimalloc-allocator"))]
info!(" Using system allocator (consider --features mimalloc-allocator for 10-25% speedup)");
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 54-feature extraction pipeline
// Features 0-4: OHLCV (normalized)
// Features 5-14: Technical indicators (10)
// Remaining features from extraction pipeline
info!("\n🏗️ Extracting 54-dimensional feature vectors...");
// Convert RealDataLoader bars to OHLCVBar format for feature extraction
let ohlcv_bars: Vec<OHLCVBar> = 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 54-dimensional feature vectors (requires 50-bar warmup)
let feature_vectors =
extract_ml_features(&ohlcv_bars).context("Failed to extract 54-dimensional features")?;
info!(
"✅ Extracted {} feature vectors (dim=54, warmup bars skipped=50)",
feature_vectors.len()
);
// Convert FeatureVector ([f64; 54]) to Vec<Vec<f32>> for PPO trainer
let state_dim = 54;
let market_data: Vec<Vec<f32>> = 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
None, // Single environment (standard mode)
)
.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::<f32>() / 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(())
}