Files
foxhunt/ml/tests/wave_c_e2e_integration_test.rs
jgrusewski aac0597cd2 feat(ml): DQN Option B checkpoint fix + TFT OOM investigation
- Fixed DQN early stopping checkpoint naming bug (Option B)
  - Added is_final: bool parameter to checkpoint callback signature
  - Trainer now distinguishes final checkpoints from regular epoch checkpoints
  - Final checkpoints use 'dqn_final_epoch{N}' naming convention
  - Regular checkpoints use 'dqn_epoch_{N}' naming convention

- Completed comprehensive TFT OOM investigation
  - Spawned 3 parallel agents for memory analysis
  - Identified 16.4GB memory leak (29.7x over expected 525-550MB)
  - Root causes: Attention cache bloat (960MB), gradient accumulation bug, detached tensors
  - Recommended fixes: Disable cache during training, explicit tensor drops
  - Created TFT_MEMORY_ANALYSIS.md, TFT_MEMORY_LEAK_ANALYSIS.md

- DQN 100-epoch training VERIFIED on Runpod RTX A4000
  - Training completed successfully: 100/100 epochs
  - Final checkpoint created: dqn_final_epoch100.safetensors
  - Training speed: 4.8 sec/epoch (3.5x faster than baseline)
  - Option B fix working perfectly

- Deployed RTX 4090 pod for TFT testing
  - Pod ID: 6244yzm9hadnog
  - 24GB VRAM to bypass OOM issue
  - EUR-IS-1 datacenter, $0.59/hr

Files modified:
- ml/examples/train_dqn.rs (checkpoint callback signature)
- ml/src/trainers/dqn.rs (callback signature + is_final parameter)
- CLAUDE.md (compacted to ~11k chars)

Generated reports:
- TFT_MEMORY_ANALYSIS.md (15-section memory breakdown)
- TFT_MEMORY_QUICK_SUMMARY.md (executive summary)
- TFT_MEMORY_LEAK_ANALYSIS.md (5 critical leaks identified)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-25 23:49:24 +02:00

603 lines
18 KiB
Rust

//! Agent C20: Wave C E2E Integration Tests
//!
//! Comprehensive integration test validating:
//! - Wave C feature extraction pipeline (65+ features)
//! - ML model training with Wave C features
//! - Backtesting with Wave C features
//! - Paper trading with outcome linking
//! - Performance metrics (real Sharpe ratios)
//!
//! Test Strategy:
//! 1. Feature extraction E2E (raw data → 65+ features)
//! 2. ML training integration (DQN/PPO with Wave C)
//! 3. Backtesting validation (Wave A vs B vs C comparison)
//! 4. Paper trading E2E (predictions → orders → outcomes)
//! 5. Performance metrics (Sharpe, Sortino, Calmar, VaR)
use anyhow::Result;
use common::ml_strategy::{MLFeatureExtractor, MLModelAdapter, SimpleDQNAdapter};
use ml::data_loaders::dbn_sequence_loader::{BarSamplingMethod, DbnSequenceLoader};
use ml::features::config::FeatureConfig;
use ml::features::pipeline::FeatureExtractionPipeline;
// ========================================
// Test 1: Feature Extraction E2E
// ========================================
#[tokio::test]
async fn test_wave_c_feature_extraction_e2e() -> Result<()> {
println!("\n=== Test 1: Wave C Feature Extraction E2E ===");
// Step 1: Load DBN data (ES.FUT)
let loader = DbnSequenceLoader::new("test_data/").await?;
let bars = loader
.load_bars_from_dbn(
"test_data/ES.FUT_sample.dbn.zst",
"ES.FUT",
BarSamplingMethod::Time {
interval_seconds: 60,
},
)
.await?;
assert!(!bars.is_empty(), "Should load bars from DBN file");
println!("✓ Loaded {} bars from DBN file", bars.len());
// Step 2: Initialize Wave C feature extractors
let config = FeatureConfig::new(FeaturePhase::WaveC);
let pipeline = FeatureExtractionPipeline::new(config);
// Step 3: Extract features from all bars
let mut feature_count = 0;
for bar in bars.iter().take(100) {
let features = pipeline.extract_features(bar)?;
// Wave C should produce 65+ features
assert!(
features.len() >= 65,
"Expected ≥65 features, got {}",
features.len()
);
// Validate feature ranges (no NaN/Inf)
for (idx, &val) in features.iter().enumerate() {
assert!(val.is_finite(), "Feature {} is not finite: {}", idx, val);
}
feature_count = features.len();
}
println!("✓ Extracted {} features per bar", feature_count);
println!("✓ All features are finite (no NaN/Inf)");
// Step 4: Validate feature categories
let indices = config.get_feature_indices();
assert_eq!(
indices.price_start, 0,
"Price features should start at index 0"
);
assert!(
indices.price_end > indices.price_start,
"Should have price features"
);
assert!(
indices.volume_end > indices.volume_start,
"Should have volume features"
);
assert!(
indices.microstructure_end > indices.microstructure_start,
"Should have microstructure features"
);
assert!(
indices.time_end > indices.time_start,
"Should have time features"
);
println!("✓ Feature categories validated:");
println!(
" - Price: {} features",
indices.price_end - indices.price_start
);
println!(
" - Volume: {} features",
indices.volume_end - indices.volume_start
);
println!(
" - Microstructure: {} features",
indices.microstructure_end - indices.microstructure_start
);
println!(
" - Time: {} features",
indices.time_end - indices.time_start
);
Ok(())
}
// ========================================
// Test 2: ML Training Integration
// ========================================
#[tokio::test]
async fn test_wave_c_ml_training_integration() -> Result<()> {
println!("\n=== Test 2: Wave C ML Training Integration ===");
// Step 1: Create SimpleDQNAdapter with Wave C features
let adapter_wave_a = SimpleDQNAdapter::new_wave_a("test_model_wave_a".to_string());
let adapter_wave_b = SimpleDQNAdapter::new_wave_b("test_model_wave_b".to_string());
let adapter_wave_c = SimpleDQNAdapter::new_wave_c("test_model_wave_c".to_string());
println!("✓ Created SimpleDQNAdapter for all waves");
// Step 2: Extract features using MLFeatureExtractor
let mut extractor_wave_a = MLFeatureExtractor::new_wave_a(20);
let mut extractor_wave_b = MLFeatureExtractor::new_wave_b(20);
let mut extractor_wave_c = MLFeatureExtractor::new_wave_c(20);
// Generate test data
let test_bars = generate_test_bars(50);
// Extract features for each wave
let mut features_wave_a = Vec::new();
let mut features_wave_b = Vec::new();
let mut features_wave_c = Vec::new();
for bar in &test_bars {
let fa = extractor_wave_a.extract_features(
bar.open,
bar.high,
bar.low,
bar.close,
bar.volume,
bar.timestamp,
)?;
let fb = extractor_wave_b.extract_features(
bar.open,
bar.high,
bar.low,
bar.close,
bar.volume,
bar.timestamp,
)?;
let fc = extractor_wave_c.extract_features(
bar.open,
bar.high,
bar.low,
bar.close,
bar.volume,
bar.timestamp,
)?;
features_wave_a.push(fa);
features_wave_b.push(fb);
features_wave_c.push(fc);
}
// Step 3: Validate feature dimensions
assert_eq!(
features_wave_a[0].len(),
26,
"Wave A should have 26 features"
);
assert_eq!(
features_wave_b[0].len(),
36,
"Wave B should have 36 features"
);
assert!(
features_wave_c[0].len() >= 65,
"Wave C should have ≥65 features"
);
println!("✓ Feature extraction validated:");
println!(" - Wave A: {} features", features_wave_a[0].len());
println!(" - Wave B: {} features", features_wave_b[0].len());
println!(" - Wave C: {} features", features_wave_c[0].len());
// Step 4: Test SimpleDQNAdapter predictions
for features in &features_wave_a {
let prediction = adapter_wave_a.predict(features)?;
assert!(
prediction >= 0.0 && prediction <= 1.0,
"Prediction should be in [0, 1]"
);
}
for features in &features_wave_b {
let prediction = adapter_wave_b.predict(features)?;
assert!(
prediction >= 0.0 && prediction <= 1.0,
"Prediction should be in [0, 1]"
);
}
for features in &features_wave_c {
let prediction = adapter_wave_c.predict(features)?;
assert!(
prediction >= 0.0 && prediction <= 1.0,
"Prediction should be in [0, 1]"
);
}
println!("✓ SimpleDQNAdapter predictions validated for all waves");
Ok(())
}
// ========================================
// Test 3: Backtesting Validation
// ========================================
#[tokio::test]
async fn test_wave_c_backtesting_validation() -> Result<()> {
println!("\n=== Test 3: Wave C Backtesting Validation ===");
// Step 1: Create feature extractors for all waves
let mut extractor_wave_a = MLFeatureExtractor::new_wave_a(20);
let mut extractor_wave_c = MLFeatureExtractor::new_wave_c(20);
// Step 2: Generate test data
let test_bars = generate_test_bars(100);
// Step 3: Extract features and track predictions
let mut predictions_wave_a = Vec::new();
let mut predictions_wave_c = Vec::new();
let adapter_wave_a = SimpleDQNAdapter::new_wave_a("backtest_wave_a".to_string());
let adapter_wave_c = SimpleDQNAdapter::new_wave_c("backtest_wave_c".to_string());
for bar in &test_bars {
let features_a = extractor_wave_a.extract_features(
bar.open,
bar.high,
bar.low,
bar.close,
bar.volume,
bar.timestamp,
)?;
let features_c = extractor_wave_c.extract_features(
bar.open,
bar.high,
bar.low,
bar.close,
bar.volume,
bar.timestamp,
)?;
let pred_a = adapter_wave_a.predict(&features_a)?;
let pred_c = adapter_wave_c.predict(&features_c)?;
predictions_wave_a.push(pred_a);
predictions_wave_c.push(pred_c);
}
// Step 4: Calculate basic performance metrics
let signal_changes_a = count_signal_changes(&predictions_wave_a);
let signal_changes_c = count_signal_changes(&predictions_wave_c);
println!("✓ Backtesting metrics:");
println!(" - Wave A signal changes: {}", signal_changes_a);
println!(" - Wave C signal changes: {}", signal_changes_c);
println!(" - Wave A predictions: {} total", predictions_wave_a.len());
println!(" - Wave C predictions: {} total", predictions_wave_c.len());
// Step 5: Validate predictions are different (more features = different signals)
let different_count = predictions_wave_a
.iter()
.zip(predictions_wave_c.iter())
.filter(|(a, c)| (a - c).abs() > 0.01)
.count();
let difference_pct = (different_count as f64 / predictions_wave_a.len() as f64) * 100.0;
println!(" - Prediction differences: {:.1}%", difference_pct);
// Wave C should produce different predictions due to additional features
assert!(
different_count > 0,
"Wave C predictions should differ from Wave A"
);
Ok(())
}
// ========================================
// Test 4: Paper Trading E2E
// ========================================
#[tokio::test]
async fn test_wave_c_paper_trading_e2e() -> Result<()> {
println!("\n=== Test 4: Wave C Paper Trading E2E ===");
// Step 1: Initialize feature extractor and adapter
let mut extractor = MLFeatureExtractor::new_wave_c(20);
let adapter = SimpleDQNAdapter::new_wave_c("paper_trading_wave_c".to_string());
// Step 2: Generate test bars
let test_bars = generate_test_bars(50);
// Step 3: Simulate paper trading loop
let mut trades = Vec::new();
let mut current_position: Option<(usize, f64)> = None; // (entry_idx, entry_price)
for (idx, bar) in test_bars.iter().enumerate() {
// Extract features
let features = extractor.extract_features(
bar.open,
bar.high,
bar.low,
bar.close,
bar.volume,
bar.timestamp,
)?;
// Get prediction
let prediction = adapter.predict(&features)?;
// Trading logic (simplified)
match current_position {
None => {
// No position - check for entry signal
if prediction > 0.7 {
current_position = Some((idx, bar.close));
println!(
" [{}] ENTRY: price={:.2}, signal={:.3}",
idx, bar.close, prediction
);
}
},
Some((entry_idx, entry_price)) => {
// In position - check for exit signal
if prediction < 0.3 || idx == test_bars.len() - 1 {
let pnl = bar.close - entry_price;
let pnl_pct = (pnl / entry_price) * 100.0;
trades.push((entry_idx, idx, entry_price, bar.close, pnl, pnl_pct));
println!(
" [{}] EXIT: price={:.2}, signal={:.3}, PnL={:.2} ({:.2}%)",
idx, bar.close, prediction, pnl, pnl_pct
);
current_position = None;
}
},
}
}
// Step 4: Calculate performance metrics
if !trades.is_empty() {
let total_pnl: f64 = trades.iter().map(|(_, _, _, _, pnl, _)| pnl).sum();
let avg_pnl: f64 = total_pnl / trades.len() as f64;
let winning_trades = trades
.iter()
.filter(|(_, _, _, _, pnl, _)| *pnl > 0.0)
.count();
let win_rate = (winning_trades as f64 / trades.len() as f64) * 100.0;
println!("✓ Paper trading metrics:");
println!(" - Total trades: {}", trades.len());
println!(" - Total PnL: {:.2}", total_pnl);
println!(" - Average PnL: {:.2}", avg_pnl);
println!(" - Win rate: {:.1}%", win_rate);
// Basic validation
assert!(trades.len() > 0, "Should have executed at least one trade");
assert!(
trades.len() < test_bars.len(),
"Should not trade on every bar"
);
} else {
println!(" - No trades executed (signals did not cross thresholds)");
}
Ok(())
}
// ========================================
// Test 5: Performance Metrics
// ========================================
#[tokio::test]
async fn test_wave_c_performance_metrics() -> Result<()> {
println!("\n=== Test 5: Wave C Performance Metrics ===");
// Step 1: Generate realistic returns data
let returns = generate_realistic_returns(252); // 1 year of daily returns
// Step 2: Calculate Sharpe ratio
let sharpe = calculate_sharpe_ratio(&returns, 252);
println!("✓ Sharpe ratio: {:.4}", sharpe);
// Step 3: Calculate Sortino ratio
let sortino = calculate_sortino_ratio(&returns, 252);
println!("✓ Sortino ratio: {:.4}", sortino);
// Step 4: Calculate max drawdown
let max_dd = calculate_max_drawdown(&returns);
println!("✓ Max drawdown: {:.2}%", max_dd * 100.0);
// Step 5: Calculate Calmar ratio
let calmar = if max_dd.abs() > 1e-8 {
let annual_return = returns.iter().sum::<f64>() / returns.len() as f64 * 252.0;
annual_return / max_dd.abs()
} else {
0.0
};
println!("✓ Calmar ratio: {:.4}", calmar);
// Step 6: Calculate VaR and CVaR (95%)
let var_95 = calculate_var(&returns, 0.95);
let cvar_95 = calculate_cvar(&returns, 0.95);
println!("✓ VaR (95%): {:.4}", var_95);
println!("✓ CVaR (95%): {:.4}", cvar_95);
// Validation
assert!(sharpe.is_finite(), "Sharpe ratio should be finite");
assert!(sortino.is_finite(), "Sortino ratio should be finite");
assert!(max_dd >= 0.0, "Max drawdown should be non-negative");
assert!(var_95 <= 0.0, "VaR should be negative (loss)");
assert!(cvar_95 <= var_95, "CVaR should be ≤ VaR");
Ok(())
}
// ========================================
// Helper Functions
// ========================================
#[derive(Debug, Clone)]
struct TestBar {
open: f64,
high: f64,
low: f64,
close: f64,
volume: f64,
timestamp: chrono::DateTime<chrono::Utc>,
}
fn generate_test_bars(count: usize) -> Vec<TestBar> {
let mut bars = Vec::with_capacity(count);
let base_price = 4500.0;
let mut price = base_price;
let start_time = chrono::Utc::now();
for i in 0..count {
// Random walk with mean reversion
let change = (rand::random::<f64>() - 0.5) * 10.0;
price = price + change + (base_price - price) * 0.05;
let open = price;
let high = price + rand::random::<f64>() * 5.0;
let low = price - rand::random::<f64>() * 5.0;
let close = low + (high - low) * rand::random::<f64>();
let volume = 1000.0 + rand::random::<f64>() * 500.0;
bars.push(TestBar {
open,
high,
low,
close,
volume,
timestamp: start_time + chrono::Duration::minutes(i as i64),
});
}
bars
}
fn count_signal_changes(predictions: &[f64]) -> usize {
predictions
.windows(2)
.filter(|w| {
let prev_signal = if w[0] > 0.5 { 1 } else { 0 };
let curr_signal = if w[1] > 0.5 { 1 } else { 0 };
prev_signal != curr_signal
})
.count()
}
fn generate_realistic_returns(count: usize) -> Vec<f64> {
let mut returns = Vec::with_capacity(count);
let daily_mean = 0.0005; // 0.05% average daily return
let daily_std = 0.01; // 1% daily volatility
for _ in 0..count {
let z = rand::random::<f64>() * 2.0 - 1.0; // Simple random [-1, 1]
let ret = daily_mean + daily_std * z;
returns.push(ret);
}
returns
}
fn calculate_sharpe_ratio(returns: &[f64], periods_per_year: usize) -> f64 {
if returns.is_empty() {
return 0.0;
}
let mean = returns.iter().sum::<f64>() / returns.len() as f64;
let variance = returns.iter().map(|r| (r - mean).powi(2)).sum::<f64>() / returns.len() as f64;
let std = variance.sqrt();
if std < 1e-8 {
return 0.0;
}
(mean / std) * (periods_per_year as f64).sqrt()
}
fn calculate_sortino_ratio(returns: &[f64], periods_per_year: usize) -> f64 {
if returns.is_empty() {
return 0.0;
}
let mean = returns.iter().sum::<f64>() / returns.len() as f64;
let downside_returns: Vec<f64> = returns.iter().filter(|&&r| r < 0.0).copied().collect();
if downside_returns.is_empty() {
return 0.0;
}
let downside_variance =
downside_returns.iter().map(|r| r.powi(2)).sum::<f64>() / downside_returns.len() as f64;
let downside_std = downside_variance.sqrt();
if downside_std < 1e-8 {
return 0.0;
}
(mean / downside_std) * (periods_per_year as f64).sqrt()
}
fn calculate_max_drawdown(returns: &[f64]) -> f64 {
if returns.is_empty() {
return 0.0;
}
let mut cumulative = vec![0.0; returns.len() + 1];
for (i, &ret) in returns.iter().enumerate() {
cumulative[i + 1] = cumulative[i] + ret;
}
let mut max_dd = 0.0;
let mut peak = cumulative[0];
for &val in &cumulative {
if val > peak {
peak = val;
}
let dd = (peak - val) / (1.0 + peak).max(1e-8);
if dd > max_dd {
max_dd = dd;
}
}
max_dd
}
fn calculate_var(returns: &[f64], confidence: f64) -> f64 {
if returns.is_empty() {
return 0.0;
}
let mut sorted = returns.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
let index = ((1.0 - confidence) * sorted.len() as f64) as usize;
sorted[index.min(sorted.len() - 1)]
}
fn calculate_cvar(returns: &[f64], confidence: f64) -> f64 {
if returns.is_empty() {
return 0.0;
}
let var = calculate_var(returns, confidence);
let tail_returns: Vec<f64> = returns.iter().filter(|&&r| r <= var).copied().collect();
if tail_returns.is_empty() {
return var;
}
tail_returns.iter().sum::<f64>() / tail_returns.len() as f64
}