Files
foxhunt/ml/examples/wave_comparison_full_features.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

569 lines
18 KiB
Rust

//! Wave Comparison Backtest with Full Feature Extraction
//!
//! This backtest compares Wave C (201 features) vs Wave D (54 features)
//! using the ACTUAL feature extraction from ml::features::extraction.
//!
//! Unlike the common::ml_strategy simplified extractor, this uses:
//! - ml::features::extraction::extract_ml_features() for full 54-feature extraction
//! - Real fractional differentiation features (indices 39-200)
//! - Real Wave D regime detection features (indices 201-224)
//!
//! Usage:
//! cargo run -p ml --example wave_comparison_full_features --release
use anyhow::Result;
use chrono::{DateTime, Utc};
use data::providers::databento::dbn_parser::{DbnParser, ProcessedMessage};
use ml::features::extraction::extract_ml_features;
use ml::features::FeaturePhase;
use num_traits::ToPrimitive;
use std::path::PathBuf;
/// Performance metrics for backtest
#[derive(Debug, Clone)]
struct PerformanceMetrics {
total_trades: usize,
winning_trades: usize,
win_rate: f64,
total_pnl: f64,
total_return: f64,
sharpe_ratio: f64,
max_drawdown: f64,
calmar_ratio: f64,
profit_factor: f64,
}
/// Trade record
#[derive(Debug, Clone)]
struct Trade {
entry_time: DateTime<Utc>,
exit_time: DateTime<Utc>,
entry_price: f64,
exit_price: f64,
side: TradeSide,
pnl: f64,
}
#[derive(Debug, Clone, Copy)]
enum TradeSide {
Long,
Short,
}
/// Market data bar
#[derive(Debug, Clone)]
struct MarketBar {
timestamp: DateTime<Utc>,
open: f64,
high: f64,
low: f64,
close: f64,
volume: f64,
}
/// Load market data from DBN file
fn load_market_data(dbn_path: &PathBuf) -> Result<Vec<MarketBar>> {
println!("📖 Loading market data from: {}", dbn_path.display());
let parser =
DbnParser::new().map_err(|e| anyhow::anyhow!("Failed to create DBN parser: {}", e))?;
let dbn_bytes = std::fs::read(dbn_path)?;
let messages = parser
.parse_batch(&dbn_bytes)
.map_err(|e| anyhow::anyhow!("Failed to parse DBN file: {}", e))?;
let mut bars = Vec::new();
for msg in messages {
if let ProcessedMessage::Ohlcv {
symbol: _,
timestamp,
open,
high,
low,
close,
volume,
} = msg
{
let ts_secs = (timestamp.as_nanos() / 1_000_000_000) as i64;
bars.push(MarketBar {
timestamp: DateTime::from_timestamp(ts_secs, 0).unwrap_or_else(|| Utc::now()),
open: open.to_f64(),
high: high.to_f64(),
low: low.to_f64(),
close: close.to_f64(),
volume: volume.to_f64().unwrap_or(0.0),
});
}
}
bars.sort_by_key(|bar| bar.timestamp);
println!("✅ Loaded {} bars", bars.len());
Ok(bars)
}
/// Calculate maximum drawdown from equity curve
fn calculate_max_drawdown(equity_curve: &[f64]) -> f64 {
if equity_curve.is_empty() {
return 0.0;
}
let mut max_drawdown = 0.0;
let mut peak = equity_curve[0];
for &equity in equity_curve {
if equity > peak {
peak = equity;
}
let drawdown = (peak - equity) / peak;
if drawdown > max_drawdown {
max_drawdown = drawdown;
}
}
max_drawdown
}
/// Extract momentum signal from feature vector
/// Uses features 5-10 (technical indicators) + Wave D regime features (201-224)
fn extract_momentum_signal(features: &[f64], use_regime_features: bool) -> f64 {
if features.len() < 10 {
return 0.0;
}
// Base signal from technical indicators (features 5-10)
let signal1 = features.get(5).cloned().unwrap_or(0.0);
let signal2 = features.get(6).cloned().unwrap_or(0.0);
let signal3 = features.get(7).cloned().unwrap_or(0.0);
let signal4 = features.get(8).cloned().unwrap_or(0.0);
let signal5 = features.get(9).cloned().unwrap_or(0.0);
let base_signal = (signal1 + signal2 + signal3 + signal4 + signal5) / 5.0;
if !use_regime_features || features.len() < 54 {
return base_signal.clamp(-1.0, 1.0);
}
// Wave D: Add regime detection signal (features 201-224)
// CUSUM statistics (201-210): structural break detection
let cusum_signal = (201..=210)
.filter_map(|i| features.get(i).cloned())
.sum::<f64>()
/ 10.0;
// ADX & Directional (211-215): trend strength
let adx_signal = (211..=215)
.filter_map(|i| features.get(i).cloned())
.sum::<f64>()
/ 5.0;
// Transition probabilities (216-220): regime persistence
let transition_signal = (216..=220)
.filter_map(|i| features.get(i).cloned())
.sum::<f64>()
/ 5.0;
// Adaptive metrics (221-224): dynamic strategy adjustment
let adaptive_signal = (221..=224)
.filter_map(|i| features.get(i).cloned())
.sum::<f64>()
/ 4.0;
// Combine signals: 60% base + 10% each regime component
let combined_signal = 0.6 * base_signal
+ 0.1 * cusum_signal
+ 0.1 * adx_signal
+ 0.1 * transition_signal
+ 0.1 * adaptive_signal;
combined_signal.clamp(-1.0, 1.0)
}
/// Run momentum-based backtest with full feature extraction
fn run_backtest(
market_data: &[MarketBar],
initial_capital: f64,
feature_phase: FeaturePhase,
wave_name: &str,
) -> Result<PerformanceMetrics> {
println!("\n🔄 Running {} backtest...", wave_name);
let mut trades = Vec::new();
let mut position: Option<(TradeSide, f64, DateTime<Utc>, f64)> = None;
let mut equity_curve = vec![initial_capital];
let mut current_capital = initial_capital;
// Simple momentum strategy parameters
let signal_threshold = 0.15;
let holding_periods = 20;
let use_regime_features = matches!(feature_phase, FeaturePhase::WaveD);
let mut bars_in_position = 0;
let mut feature_buffer: Vec<Vec<f64>> = Vec::new();
for bar in market_data.iter() {
// Extract features using ml::features::extraction
let features = extract_ml_features(
bar.open,
bar.high,
bar.low,
bar.close,
bar.volume,
feature_phase,
);
// For Wave C, zero out features 201-224 to simulate pure Wave C performance
let filtered_features = if matches!(feature_phase, FeaturePhase::WaveC) {
let mut f = features;
// Zero out Wave D features
if f.len() >= 54 {
for i in 201..224 {
f[i] = 0.0;
}
}
f
} else {
features
};
feature_buffer.push(filtered_features.clone());
// Keep only last 10 bars for lookback
if feature_buffer.len() > 10 {
feature_buffer.remove(0);
}
// Get momentum signal
let signal = extract_momentum_signal(&filtered_features, use_regime_features);
// Trading logic
if position.is_none() && signal.abs() > signal_threshold {
// Enter position
let side = if signal > 0.0 {
TradeSide::Long
} else {
TradeSide::Short
};
let size = (current_capital * 0.1) / bar.close;
position = Some((side, size, bar.timestamp, bar.close));
bars_in_position = 0;
} else if let Some((side, size, entry_time, entry_price)) = position {
bars_in_position += 1;
// Exit logic
let should_exit = match side {
TradeSide::Long => signal < -0.1 || bars_in_position >= holding_periods,
TradeSide::Short => signal > 0.1 || bars_in_position >= holding_periods,
};
if should_exit {
let pnl = match side {
TradeSide::Long => size * (bar.close - entry_price),
TradeSide::Short => size * (entry_price - bar.close),
};
current_capital += pnl;
equity_curve.push(current_capital);
trades.push(Trade {
entry_time,
exit_time: bar.timestamp,
entry_price,
exit_price: bar.close,
side,
pnl,
});
position = None;
bars_in_position = 0;
}
}
}
// Close any open position
if let Some((side, size, entry_time, entry_price)) = position {
let last_bar = &market_data[market_data.len() - 1];
let pnl = match side {
TradeSide::Long => size * (last_bar.close - entry_price),
TradeSide::Short => size * (entry_price - last_bar.close),
};
current_capital += pnl;
equity_curve.push(current_capital);
trades.push(Trade {
entry_time,
exit_time: last_bar.timestamp,
entry_price,
exit_price: last_bar.close,
side,
pnl,
});
}
// Calculate metrics
let total_trades = trades.len();
let winning_trades = trades.iter().filter(|t| t.pnl > 0.0).count();
let win_rate = if total_trades > 0 {
(winning_trades as f64 / total_trades as f64) * 100.0
} else {
0.0
};
let total_pnl: f64 = trades.iter().map(|t| t.pnl).sum();
let total_return = (current_capital - initial_capital) / initial_capital * 100.0;
// Sharpe ratio (annualized)
let returns: Vec<f64> = trades.iter().map(|t| t.pnl / initial_capital).collect();
let sharpe_ratio = if !returns.is_empty() {
let mean_return = returns.iter().sum::<f64>() / returns.len() as f64;
let variance = returns
.iter()
.map(|r| (r - mean_return).powi(2))
.sum::<f64>()
/ returns.len() as f64;
let std_dev = variance.sqrt();
if std_dev > 0.0 {
(mean_return / std_dev) * (252.0_f64).sqrt()
} else {
0.0
}
} else {
0.0
};
let max_drawdown = calculate_max_drawdown(&equity_curve) * 100.0;
let calmar_ratio = if max_drawdown > 0.0 {
total_return / max_drawdown
} else {
0.0
};
let gross_profit: f64 = trades.iter().filter(|t| t.pnl > 0.0).map(|t| t.pnl).sum();
let gross_loss: f64 = trades
.iter()
.filter(|t| t.pnl < 0.0)
.map(|t| t.pnl.abs())
.sum();
let profit_factor = if gross_loss > 0.0 {
gross_profit / gross_loss
} else if gross_profit > 0.0 {
f64::INFINITY
} else {
0.0
};
Ok(PerformanceMetrics {
total_trades,
winning_trades,
win_rate,
total_pnl,
total_return,
sharpe_ratio,
max_drawdown,
calmar_ratio,
profit_factor,
})
}
fn print_metrics(metrics: &PerformanceMetrics, wave_name: &str) {
println!("\n{}", "=".repeat(70));
println!("📈 {} RESULTS", wave_name.to_uppercase());
println!("{}", "=".repeat(70));
println!("\n💰 Performance Metrics:");
println!(" Total Trades: {}", metrics.total_trades);
println!(" Winning Trades: {}", metrics.winning_trades);
println!(" Win Rate: {:.2}%", metrics.win_rate);
println!(" Total PnL: ${:.2}", metrics.total_pnl);
println!(" Total Return: {:.2}%", metrics.total_return);
println!(" Sharpe Ratio: {:.2}", metrics.sharpe_ratio);
println!(" Max Drawdown: {:.2}%", metrics.max_drawdown);
println!(" Calmar Ratio: {:.2}", metrics.calmar_ratio);
println!(" Profit Factor: {:.2}", metrics.profit_factor);
}
fn print_comparison(wave_c: &PerformanceMetrics, wave_d: &PerformanceMetrics) {
println!("\n{}", "=".repeat(70));
println!("📊 WAVE C vs WAVE D COMPARISON");
println!("{}", "=".repeat(70));
let sharpe_improvement =
((wave_d.sharpe_ratio - wave_c.sharpe_ratio) / wave_c.sharpe_ratio.abs().max(0.01)) * 100.0;
let win_rate_improvement = wave_d.win_rate - wave_c.win_rate;
let drawdown_improvement =
((wave_c.max_drawdown - wave_d.max_drawdown) / wave_c.max_drawdown.abs().max(0.01)) * 100.0;
let return_improvement = wave_d.total_return - wave_c.total_return;
println!("\n🎯 Key Improvements:");
println!(
" Sharpe Ratio: {:.2}{:.2} ({:+.1}%)",
wave_c.sharpe_ratio, wave_d.sharpe_ratio, sharpe_improvement
);
println!(
" Win Rate: {:.2}% → {:.2}% ({:+.1}pp)",
wave_c.win_rate, wave_d.win_rate, win_rate_improvement
);
println!(
" Max Drawdown: {:.2}% → {:.2}% ({:+.1}%)",
wave_c.max_drawdown, wave_d.max_drawdown, drawdown_improvement
);
println!(
" Total Return: {:.2}% → {:.2}% ({:+.2}pp)",
wave_c.total_return, wave_d.total_return, return_improvement
);
println!("\n✅ Target Validation:");
println!(
" Sharpe ≥ 2.0: {} (actual: {:.2})",
if wave_d.sharpe_ratio >= 2.0 {
"✅ PASS"
} else {
"❌ FAIL"
},
wave_d.sharpe_ratio
);
println!(
" Win Rate ≥ 60%: {} (actual: {:.2}%)",
if wave_d.win_rate >= 60.0 {
"✅ PASS"
} else {
"❌ FAIL"
},
wave_d.win_rate
);
println!(
" Drawdown ≤ 15%: {} (actual: {:.2}%)",
if wave_d.max_drawdown <= 15.0 {
"✅ PASS"
} else {
"❌ FAIL"
},
wave_d.max_drawdown
);
let all_targets_met =
wave_d.sharpe_ratio >= 2.0 && wave_d.win_rate >= 60.0 && wave_d.max_drawdown <= 15.0;
println!(
"\n{}",
if all_targets_met {
"🎉 ALL TARGETS MET - PRODUCTION READY!"
} else {
"⚠️ Some targets not met - further optimization needed"
}
);
}
fn main() -> Result<()> {
println!("\n{}", "=".repeat(70));
println!("🚀 WAVE COMPARISON BACKTEST (Full Feature Extraction)");
println!("{}\n", "=".repeat(70));
let data_path = PathBuf::from(
"/home/jgrusewski/Work/foxhunt/test_data/real/databento/ES.FUT_ohlcv-1m_2024-01-02.dbn",
);
let initial_capital = 100000.0;
// Load market data
let market_data = load_market_data(&data_path)?;
println!("📊 Backtest Configuration:");
println!(" Symbol: ES.FUT");
println!(" Bars: {}", market_data.len());
println!(" Initial Capital: ${:.2}", initial_capital);
println!(" Strategy: Momentum + Regime Detection");
println!(" Feature Extraction: ml::features::extraction (FULL 54 features)");
// Run Wave C backtest (201 features, no regime)
let wave_c_metrics = run_backtest(
&market_data,
initial_capital,
FeaturePhase::WaveC,
"Wave C (201 features, no regime)",
)?;
print_metrics(&wave_c_metrics, "Wave C (201 features)");
// Run Wave D backtest (54 features, with regime)
let wave_d_metrics = run_backtest(
&market_data,
initial_capital,
FeaturePhase::WaveD,
"Wave D (54 features, with regime)",
)?;
print_metrics(&wave_d_metrics, "Wave D (54 features)");
// Print comparison
print_comparison(&wave_c_metrics, &wave_d_metrics);
println!("\n{}", "=".repeat(70));
// Save results
let report = format!(
"# Wave Comparison Backtest Results (Full Feature Extraction)\n\n\
## Wave C (201 Features)\n\
- Total Trades: {}\n\
- Win Rate: {:.2}%\n\
- Total Return: {:.2}%\n\
- Sharpe Ratio: {:.2}\n\
- Max Drawdown: {:.2}%\n\n\
## Wave D (54 Features + Regime Detection)\n\
- Total Trades: {}\n\
- Win Rate: {:.2}%\n\
- Total Return: {:.2}%\n\
- Sharpe Ratio: {:.2}\n\
- Max Drawdown: {:.2}%\n\n\
## Improvements\n\
- Sharpe: {:.2}{:.2} ({:+.1}%)\n\
- Win Rate: {:.2}% → {:.2}% ({:+.1}pp)\n\
- Drawdown: {:.2}% → {:.2}% ({:+.1}%)\n\n\
## Target Validation\n\
- Sharpe ≥ 2.0: {}\n\
- Win Rate ≥ 60%: {}\n\
- Drawdown ≤ 15%: {}\n",
wave_c_metrics.total_trades,
wave_c_metrics.win_rate,
wave_c_metrics.total_return,
wave_c_metrics.sharpe_ratio,
wave_c_metrics.max_drawdown,
wave_d_metrics.total_trades,
wave_d_metrics.win_rate,
wave_d_metrics.total_return,
wave_d_metrics.sharpe_ratio,
wave_d_metrics.max_drawdown,
wave_c_metrics.sharpe_ratio,
wave_d_metrics.sharpe_ratio,
((wave_d_metrics.sharpe_ratio - wave_c_metrics.sharpe_ratio)
/ wave_c_metrics.sharpe_ratio.abs().max(0.01))
* 100.0,
wave_c_metrics.win_rate,
wave_d_metrics.win_rate,
wave_d_metrics.win_rate - wave_c_metrics.win_rate,
wave_c_metrics.max_drawdown,
wave_d_metrics.max_drawdown,
((wave_c_metrics.max_drawdown - wave_d_metrics.max_drawdown)
/ wave_c_metrics.max_drawdown.abs().max(0.01))
* 100.0,
if wave_d_metrics.sharpe_ratio >= 2.0 {
"PASS ✅"
} else {
"FAIL ❌"
},
if wave_d_metrics.win_rate >= 60.0 {
"PASS ✅"
} else {
"FAIL ❌"
},
if wave_d_metrics.max_drawdown <= 15.0 {
"PASS ✅"
} else {
"FAIL ❌"
},
);
std::fs::write("/tmp/wave_comparison_backtest_v2.md", report)?;
println!("\n💾 Results saved to: /tmp/wave_comparison_backtest_v2.md");
Ok(())
}