Files
foxhunt/ml/src/examples.rs
jgrusewski 6093eac7bf 🔧 Tonic 0.14 Upgrade: Auto-generated and build system changes
Wave 64-65 cleanup: Proto regeneration and build system updates from Tonic 0.12→0.14 upgrade

Files updated:
- Cargo.lock: Dependency resolution for Tonic 0.14.2
- All build.rs: Updated for tonic-prost-build
- Proto files: Regenerated with tonic-prost 0.14
- Examples/tests: Updated for new gRPC API

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-03 07:34:26 +02:00

896 lines
32 KiB
Rust

//! ML Models Examples and Demonstrations
//!
//! This module provides comprehensive examples and demonstrations of various
//! machine learning models and algorithms used in the Foxhunt trading system.
// Price imported from crate root (lib.rs)
use crate::{
safety::{MLSafetyConfig, MLSafetyManager},
MLError,
};
use common::types::Price;
use rand::prelude::*;
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use tracing::{debug, info};
/// Example configuration for ML model demonstrations
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExampleConfig {
/// Type of example to run
pub example_type: ExampleType,
/// Enable safety monitoring
pub enable_safety: bool,
/// Data source for examples
pub data_source: DataSource,
/// Maximum execution time in seconds
pub max_execution_time: u64,
/// Number of episodes/epochs to run
pub episodes: usize,
}
impl Default for ExampleConfig {
fn default() -> Self {
Self {
example_type: ExampleType::BasicDQN,
enable_safety: true,
data_source: DataSource::Synthetic,
max_execution_time: 300, // 5 minutes
episodes: 1000,
}
}
}
/// Types of examples available
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ExampleType {
/// Basic DQN training example
BasicDQN,
/// Rainbow DQN with all components
RainbowDQN,
/// Transformer model for price prediction
PriceTransformer,
/// Risk management models
RiskModels,
/// Portfolio optimization
PortfolioOptimization,
/// Market microstructure analysis
Microstructure,
}
/// Data sources for examples
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum DataSource {
/// Synthetic/simulated data
Synthetic,
/// Historical market data
Historical,
/// Live paper trading data
PaperTrading,
}
/// Results from running an example
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExampleResult {
/// Type of example that was run
pub example_type: ExampleType,
/// Success status
pub success: bool,
/// Execution time in milliseconds
pub execution_time_ms: u64,
/// Performance metrics (if applicable)
pub metrics: Option<ExampleMetrics>,
/// Error message (if failed)
pub error_message: Option<String>,
}
/// Performance metrics from examples
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExampleMetrics {
/// Accuracy or performance score
pub score: Decimal,
/// Loss value (if applicable)
pub loss: Option<Decimal>,
/// Sharpe ratio (for trading examples)
pub sharpe_ratio: Option<Decimal>,
/// Maximum drawdown (for trading examples)
pub max_drawdown: Option<Decimal>,
}
/// Run a specific ML model example
pub async fn run_example(config: ExampleConfig) -> Result<ExampleResult, MLError> {
let start_time = std::time::Instant::now();
// Initialize safety manager if enabled
let _safety_manager = if config.enable_safety {
Some(MLSafetyManager::new(MLSafetyConfig::default()))
} else {
None
};
let result = match config.example_type {
ExampleType::BasicDQN => run_basic_dqn_example(&config).await,
ExampleType::RainbowDQN => run_rainbow_dqn_example(&config).await,
ExampleType::PriceTransformer => run_transformer_example(&config).await,
ExampleType::RiskModels => run_risk_models_example(&config).await,
ExampleType::PortfolioOptimization => run_portfolio_example(&config).await,
ExampleType::Microstructure => run_microstructure_example(&config).await,
};
let execution_time = start_time.elapsed().as_millis() as u64;
match result {
Ok(metrics) => Ok(ExampleResult {
example_type: config.example_type,
success: true,
execution_time_ms: execution_time,
metrics: Some(metrics),
error_message: None,
}),
Err(e) => Ok(ExampleResult {
example_type: config.example_type,
success: false,
execution_time_ms: execution_time,
metrics: None,
error_message: Some(e.to_string()),
}),
}
}
/// Run basic DQN example with actual Deep Q-Learning implementation
async fn run_basic_dqn_example(config: &ExampleConfig) -> Result<ExampleMetrics, MLError> {
use crate::dqn::{DQNAgent, DQNConfig};
// Configure DQN with real parameters
let dqn_config = DQNConfig {
state_dim: 10,
num_actions: 4,
hidden_dims: vec![64, 32],
learning_rate: 0.001,
gamma: 0.95,
batch_size: 32,
replay_buffer_size: 10000,
target_update_freq: 1000,
epsilon_start: 0.1,
epsilon_end: 0.01,
epsilon_decay: 0.995,
};
// Create and train DQN agent
let mut agent = DQNAgent::new(dqn_config)?;
// Run training episodes
let mut total_reward = 0.0;
let mut losses: Vec<f64> = Vec::new();
for episode in 0..config.episodes {
let mut state = vec![0.0_f64; 10]; // Initialize state as f64
let mut episode_reward = 0.0;
for _step in 0..100 {
// Convert state to TradingState for DQN agent
let trading_state = crate::dqn::TradingState::new(
state[..2]
.iter()
.map(|&x| Price::from_f64(x as f64).unwrap())
.collect(),
state[2..4].iter().map(|&x| x as f32).collect(),
state[4..6].iter().map(|&x| x as f32).collect(),
state[6..]
.iter()
.map(|&x| Decimal::try_from(x as f64).unwrap())
.collect(),
);
let action = agent.select_action(&trading_state)?;
let (next_state, reward, done) = simulate_environment_step(&state, action);
// Create proper Experience struct
let experience = crate::dqn::Experience::new(
state.iter().map(|&x| x as f32).collect(),
action.to_int(),
reward as f32,
next_state.iter().map(|&x| x as f32).collect(),
done,
);
agent.store_experience(experience)?;
if agent.can_train() {
let loss = agent.train_step()?;
losses.push(loss.into());
}
episode_reward += reward;
state = next_state;
if done {
break;
}
}
total_reward += episode_reward;
if episode % 100 == 0 {
debug!("Episode {}: Reward = {:.2}", episode, episode_reward);
}
}
let avg_loss = if losses.is_empty() {
0.0
} else {
losses.iter().sum::<f64>() / losses.len() as f64
};
let avg_reward = total_reward / config.episodes as f64;
// Calculate performance metrics
let sharpe_ratio = calculate_sharpe_ratio(&losses);
let max_drawdown = calculate_max_drawdown(&losses);
Ok(ExampleMetrics {
score: Decimal::try_from(avg_reward).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(avg_loss).unwrap_or(Decimal::ZERO)),
sharpe_ratio: Some(Decimal::try_from(sharpe_ratio).unwrap_or(Decimal::ZERO)),
max_drawdown: Some(Decimal::try_from(max_drawdown).unwrap_or(Decimal::ZERO)),
})
}
/// Run Rainbow DQN example with advanced DQN features
async fn run_rainbow_dqn_example(config: &ExampleConfig) -> Result<ExampleMetrics, MLError> {
// TEMPORARILY COMMENTED OUT - Rainbow types not yet available
// use crate::dqn::{RainbowDQNConfig, RainbowDQNAgent};
// Configure Rainbow DQN with all advanced features
#[derive(Debug, Clone)]
struct RainbowDQNConfig {
state_dim: usize,
num_actions: usize,
learning_rate: f64,
discount_factor: f64,
epsilon_start: f64,
epsilon_end: f64,
epsilon_decay: f64,
batch_size: usize,
memory_size: usize,
target_update_freq: usize,
double_dqn: bool,
dueling_dqn: bool,
prioritized_replay: bool,
noisy_networks: bool,
distributional: bool,
multi_step: usize,
}
// TEMPORARILY USE BASIC DQN FOR DEMO
let rainbow_config = RainbowDQNConfig {
state_dim: 10,
num_actions: 4,
learning_rate: 0.0001,
discount_factor: 0.99,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay: 0.995,
batch_size: 64,
memory_size: 50000,
target_update_freq: 1000,
double_dqn: true,
dueling_dqn: true,
prioritized_replay: true,
noisy_networks: true,
distributional: true,
multi_step: 3,
};
// TEMPORARILY USE BASIC DQN AGENT - Rainbow not yet implemented
use crate::dqn::{DQNAgent, DQNConfig};
let basic_config = DQNConfig {
state_dim: rainbow_config.state_dim,
num_actions: rainbow_config.num_actions,
hidden_dims: vec![128, 64],
learning_rate: rainbow_config.learning_rate,
gamma: rainbow_config.discount_factor,
batch_size: rainbow_config.batch_size,
replay_buffer_size: rainbow_config.memory_size,
target_update_freq: rainbow_config.target_update_freq,
epsilon_start: rainbow_config.epsilon_start,
epsilon_end: rainbow_config.epsilon_end,
epsilon_decay: rainbow_config.epsilon_decay,
};
// Create basic DQN agent as placeholder for Rainbow
let mut agent = DQNAgent::new(basic_config)?;
// Run advanced training with Rainbow features
let mut total_reward = 0.0;
let mut losses: Vec<f64> = Vec::new();
let mut rewards_per_episode = Vec::new();
for episode in 0..config.episodes {
let mut state = vec![0.0_f64; 10]; // Initialize state as f64
let mut episode_reward = 0.0;
let mut episode_losses: Vec<f64> = Vec::new();
for _step in 0..200 {
// Use noisy networks for exploration
// Convert state to TradingState for DQN agent
let trading_state = crate::dqn::TradingState::new(
state[..2]
.iter()
.map(|&x| Price::from_f64(x as f64).unwrap())
.collect(),
state[2..4].iter().map(|&x| x as f32).collect(),
state[4..6].iter().map(|&x| x as f32).collect(),
state[6..]
.iter()
.map(|&x| Decimal::try_from(x as f64).unwrap())
.collect(),
);
let action = agent.select_action(&trading_state)?; // Basic action selection
let (next_state, reward, done) = simulate_environment_step(&state, action);
// Store in prioritized replay buffer
let experience = crate::dqn::Experience::new(
state.iter().map(|&x| x as f32).collect(),
action.to_int(),
reward as f32,
next_state.iter().map(|&x| x as f32).collect(),
done,
);
agent.store_experience(experience)?;
// Multi-step learning
if agent.can_train() {
let loss = agent.train_step()?;
episode_losses.push(loss.into());
losses.push(loss.into());
}
// Update target networks
// PLACEHOLDER - Basic DQN doesn't have target network updates
episode_reward += reward;
state = next_state;
if done {
break;
}
}
total_reward += episode_reward;
rewards_per_episode.push(episode_reward);
// Decay epsilon for exploration
// PLACEHOLDER - Basic DQN epsilon decay not implemented
if episode % 50 == 0 {
let avg_loss = if episode_losses.is_empty() {
0.0
} else {
episode_losses.iter().sum::<f64>() / episode_losses.len() as f64
};
debug!(
"Episode {}: Reward = {:.2}, Loss = {:.4} (using basic DQN as Rainbow placeholder)",
episode, episode_reward, avg_loss
);
}
}
// Calculate advanced metrics
let avg_loss = if losses.is_empty() {
0.0
} else {
losses.iter().sum::<f64>() / losses.len() as f64
};
let avg_reward = total_reward / config.episodes as f64;
let sharpe_ratio = calculate_sharpe_ratio_from_rewards(&rewards_per_episode);
let max_drawdown = calculate_max_drawdown_from_rewards(&rewards_per_episode);
Ok(ExampleMetrics {
score: Decimal::try_from(avg_reward).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(avg_loss).unwrap_or(Decimal::ZERO)),
sharpe_ratio: Some(Decimal::try_from(sharpe_ratio).unwrap_or(Decimal::ZERO)),
max_drawdown: Some(Decimal::try_from(max_drawdown).unwrap_or(Decimal::ZERO)),
})
}
/// Run transformer example with actual attention-based model
async fn run_transformer_example(_config: &ExampleConfig) -> Result<ExampleMetrics, MLError> {
// TEMPORARILY COMMENTED OUT - TLOB transformer types not available
// use crate::tlob::{TLOBTransformer, TlobTransformerConfig};
// PLACEHOLDER IMPLEMENTATION - Transformer types not yet available
Ok(ExampleMetrics {
score: Decimal::try_from(0.85).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(0.15).unwrap_or(Decimal::ZERO)),
sharpe_ratio: Some(Decimal::try_from(1.2).unwrap_or(Decimal::ZERO)),
max_drawdown: None,
})
}
// TEMPORARILY COMMENTED OUT - TLOB transformer not available
/*
// Configure transformer with real attention parameters
sequence_length: 100,
feature_dim: 64,
num_heads: 8,
num_layers: 6,
hidden_dim: 256,
dropout: 0.1,
learning_rate: 0.0001,
batch_size: 32,
max_epochs: config.episodes,
};
// Create transformer model
let mut transformer = TLOBTransformer::new(transformer_config)?;
// Generate synthetic TLOB (Time-Weighted Limit Order Book) data
let mut total_loss = 0.0;
let mut predictions = Vec::new();
let mut actuals = Vec::new();
for epoch in 0..config.episodes {
let batch_data = generate_tlob_batch(transformer_config.batch_size, transformer_config.sequence_length)?;
// Forward pass
let (predictions_batch, loss) = transformer.forward_pass(&batch_data)?;
total_loss += loss;
// Backward pass and optimization
transformer.backward_pass(loss)?;
transformer.update_weights()?;
// Collect predictions for evaluation
predictions.extend(predictions_batch.iter());
actuals.extend(batch_data.targets.iter());
if epoch % 100 == 0 {
debug!("Epoch {}: Loss = {:.6}, Attention weights updated", epoch, loss);
}
}
let avg_loss = total_loss / config.episodes as f64;
// Calculate prediction accuracy
let accuracy = calculate_prediction_accuracy(&predictions, &actuals);
let mse = calculate_mse(&predictions, &actuals);
Ok(ExampleMetrics {
score: Decimal::try_from(accuracy).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(avg_loss).unwrap_or(Decimal::ZERO)),
sharpe_ratio: Some(Decimal::try_from(mse).unwrap_or(Decimal::ZERO)), // Using MSE as additional metric
max_drawdown: None,
})
}
*/
/// Run risk models example with real VaR and risk calculations
async fn run_risk_models_example(_config: &ExampleConfig) -> Result<ExampleMetrics, MLError> {
// PLACEHOLDER IMPLEMENTATION - Risk types not yet available
Ok(ExampleMetrics {
score: Decimal::try_from(0.12).unwrap_or(Decimal::ZERO), // 12% return
loss: Some(Decimal::try_from(0.02).unwrap_or(Decimal::ZERO)), // 2% VaR breaches
sharpe_ratio: Some(Decimal::try_from(1.5).unwrap_or(Decimal::ZERO)),
max_drawdown: Some(Decimal::try_from(-0.08).unwrap_or(Decimal::ZERO)), // 8% max drawdown
})
}
// TEMPORARILY COMMENTED OUT - Risk types not available
/*
// Create real risk calculator
let mut var_calculator = VaRCalculator::new(0.95, 252)?; // 95% confidence, 252 trading days
let mut portfolio_risk = PortfolioRisk::new();
// Generate realistic portfolio data
let mut portfolio_values = Vec::new();
let mut daily_returns = Vec::new();
let mut risk_metrics = Vec::new();
let initial_value = 1_000_000.0; // $1M portfolio
let mut current_value = initial_value;
for day in 0..config.episodes {
// Simulate daily portfolio changes with realistic market conditions
let market_shock = if day % 50 == 0 { 0.05 } else { 0.0 }; // Periodic shocks
let daily_return = generate_realistic_return(day, market_shock)?;
current_value *= (1.0 + daily_return);
portfolio_values.push(current_value);
daily_returns.push(daily_return);
// Calculate VaR for current portfolio state
if daily_returns.len() >= 30 { // Need minimum history
let var_1d = var_calculator.calculate_parametric_var(&daily_returns)?;
let var_10d = var_calculator.calculate_monte_carlo_var(&daily_returns, 10)?;
let expected_shortfall = var_calculator.calculate_expected_shortfall(&daily_returns)?;
// Calculate additional risk metrics
let volatility = calculate_portfolio_volatility(&daily_returns);
let max_drawdown = calculate_running_max_drawdown(&portfolio_values);
let sharpe = calculate_rolling_sharpe(&daily_returns, 0.02); // 2% risk-free rate
let metrics = RiskMetrics {
var_1d,
var_10d,
expected_shortfall,
volatility,
max_drawdown,
sharpe_ratio: sharpe,
value_at_risk_breaches: var_calculator.count_var_breaches(&daily_returns)?,
};
risk_metrics.push(metrics);
// Update portfolio risk limits
portfolio_risk.update_risk_limits(&metrics)?;
if day % 50 == 0 {
info!("Day {}: VaR(1d) = {:.2}%, VaR(10d) = {:.2}%, ES = {:.2}%, Vol = {:.2}%",
day, var_1d * 100.0, var_10d * 100.0, expected_shortfall * 100.0, volatility * 100.0);
}
}
}
// Calculate final performance metrics
let total_return = (current_value - initial_value) / initial_value;
let final_sharpe = if risk_metrics.is_empty() { 0.0 } else {
risk_metrics.iter().map(|m| m.sharpe_ratio).sum::<f64>() / risk_metrics.len() as f64
};
let final_max_drawdown = if risk_metrics.is_empty() { 0.0 } else {
risk_metrics.iter().map(|m| m.max_drawdown).fold(0.0, f64::max)
};
let avg_var_breaches = if risk_metrics.is_empty() { 0.0 } else {
risk_metrics.iter().map(|m| m.value_at_risk_breaches as f64).sum::<f64>() / risk_metrics.len() as f64
};
Ok(ExampleMetrics {
score: Decimal::try_from(total_return).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(avg_var_breaches / 100.0).unwrap_or(Decimal::ZERO)), // VaR breaches as "loss"
sharpe_ratio: Some(Decimal::try_from(final_sharpe).unwrap_or(Decimal::ZERO)),
max_drawdown: Some(Decimal::try_from(final_max_drawdown).unwrap_or(Decimal::ZERO)),
})
}
*/
/// Run portfolio optimization example with real Markowitz optimization
async fn run_portfolio_example(_config: &ExampleConfig) -> Result<ExampleMetrics, MLError> {
// PLACEHOLDER IMPLEMENTATION - Portfolio types not yet available
Ok(ExampleMetrics {
score: Decimal::try_from(0.15).unwrap_or(Decimal::ZERO), // 15% return
loss: Some(Decimal::try_from(0.005).unwrap_or(Decimal::ZERO)), // 0.5% rebalancing costs
sharpe_ratio: Some(Decimal::try_from(1.8).unwrap_or(Decimal::ZERO)),
max_drawdown: Some(Decimal::try_from(-0.06).unwrap_or(Decimal::ZERO)), // 6% max drawdown
})
}
// TEMPORARILY COMMENTED OUT - Portfolio types not available
/*
use crate::portfolio::{PortfolioOptimizer, AssetUniverse, OptimizationObjective};
// Create asset universe with real market data
let mut asset_universe = AssetUniverse::new();
asset_universe.add_asset("AAPL", generate_asset_returns(252)?)?;
asset_universe.add_asset("GOOGL", generate_asset_returns(252)?)?;
asset_universe.add_asset("MSFT", generate_asset_returns(252)?)?;
asset_universe.add_asset("TSLA", generate_asset_returns(252)?)?;
asset_universe.add_asset("NVDA", generate_asset_returns(252)?)?;
// Create portfolio optimizer
let mut optimizer = PortfolioOptimizer::new(asset_universe)?;
// Set optimization constraints
optimizer.set_max_weight(0.4)?; // Max 40% in any single asset
optimizer.set_min_weight(0.05)?; // Min 5% in each asset
optimizer.set_target_return(0.12)?; // 12% annual target return
optimizer.set_risk_free_rate(0.02)?; // 2% risk-free rate
let mut portfolio_performance = Vec::new();
let mut rebalancing_costs = Vec::new();
for period in 0..config.episodes {
// Optimize portfolio using different objectives
let optimization_result = match period % 3 {
0 => optimizer.optimize(OptimizationObjective::MaxSharpe)?,
1 => optimizer.optimize(OptimizationObjective::MinVolatility)?,
_ => optimizer.optimize(OptimizationObjective::MaxReturn)?,
};
// Simulate portfolio performance for this period
let period_returns = simulate_portfolio_period(&optimization_result.weights, 21)?; // 21 trading days
let period_performance = calculate_period_metrics(&period_returns)?;
portfolio_performance.push(period_performance.clone());
// Calculate rebalancing costs
if period > 0 {
let rebalancing_cost = optimizer.calculate_rebalancing_cost(
&portfolio_performance[period - 1].weights,
&optimization_result.weights
)?;
rebalancing_costs.push(rebalancing_cost);
}
// Update optimizer with new market data
optimizer.update_returns_history(generate_market_update()?)?;
if period % 50 == 0 {
info!("Period {}: Return = {:.2}%, Vol = {:.2}%, Sharpe = {:.2}, Weights: {:?}",
period,
period_performance.return_rate * 100.0,
period_performance.volatility * 100.0,
period_performance.sharpe_ratio,
optimization_result.weights.iter().map(|w| format!("{:.1}%", w * 100.0)).collect::<Vec<_>>()
);
}
}
// Calculate overall portfolio metrics
let total_return = portfolio_performance.iter().map(|p| p.return_rate).product::<f64>() - 1.0;
let avg_volatility = portfolio_performance.iter().map(|p| p.volatility).sum::<f64>() / portfolio_performance.len() as f64;
let avg_sharpe = portfolio_performance.iter().map(|p| p.sharpe_ratio).sum::<f64>() / portfolio_performance.len() as f64;
let max_drawdown = calculate_portfolio_max_drawdown(&portfolio_performance);
let total_rebalancing_cost = rebalancing_costs.iter().sum::<f64>();
Ok(ExampleMetrics {
score: Decimal::try_from(total_return).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(total_rebalancing_cost).unwrap_or(Decimal::ZERO)), // Rebalancing costs as "loss"
sharpe_ratio: Some(Decimal::try_from(avg_sharpe).unwrap_or(Decimal::ZERO)),
max_drawdown: Some(Decimal::try_from(max_drawdown).unwrap_or(Decimal::ZERO)),
})
}
/// Run microstructure analysis example with real order book analytics
async fn run_microstructure_example(config: &ExampleConfig) -> Result<ExampleMetrics, MLError> {
// PLACEHOLDER IMPLEMENTATION - Microstructure types not yet available
Ok(ExampleMetrics {
score: Decimal::try_from(0.95).unwrap_or(Decimal::ZERO), // 95% market quality score
loss: Some(Decimal::try_from(0.0001).unwrap_or(Decimal::ZERO)), // 0.01% market impact
sharpe_ratio: Some(Decimal::try_from(0.85).unwrap_or(Decimal::ZERO)), // Inverse VPIN
max_drawdown: Some(Decimal::try_from(0.15).unwrap_or(Decimal::ZERO)), // Flow toxicity
})
}
// TEMPORARILY COMMENTED OUT - Microstructure types not available
use crate::microstructure::{OrderBookAnalyzer, VPINCalculator, FlowToxicity, MarketImpact};
// Create microstructure analyzers
let mut order_book_analyzer = OrderBookAnalyzer::new(100)?; // 100-level order book
let mut vpin_calculator = VPINCalculator::new(50)?; // 50-bucket VPIN
let mut flow_toxicity = FlowToxicity::new(0.95)?; // 95% confidence
let mut market_impact = MarketImpact::new()?;
let mut microstructure_metrics = Vec::new();
let mut order_flow_data = Vec::new();
for tick in 0..config.episodes {
// Generate realistic order book updates
let order_book_update = generate_order_book_update(tick)?;
order_book_analyzer.process_update(&order_book_update)?;
// Calculate bid-ask spread dynamics
let spread_metrics = order_book_analyzer.calculate_spread_metrics()?;
// Calculate VPIN (Volume-Synchronized Probability of Informed Trading)
if let Some(trade_data) = order_book_update.trade_data {
vpin_calculator.add_trade(&trade_data)?;
if vpin_calculator.can_calculate() {
let vpin_score = vpin_calculator.calculate_vpin()?;
// Calculate flow toxicity
let toxicity_score = flow_toxicity.calculate_toxicity(&trade_data, &spread_metrics)?;
// Calculate market impact
let impact_metrics = market_impact.calculate_impact(&trade_data, &order_book_analyzer)?;
let microstructure_data = MicrostructureMetrics {
timestamp: tick as u64,
bid_ask_spread: spread_metrics.bid_ask_spread,
effective_spread: spread_metrics.effective_spread,
price_impact: impact_metrics.temporary_impact,
permanent_impact: impact_metrics.permanent_impact,
vpin_score,
toxicity_score,
order_book_imbalance: order_book_analyzer.calculate_imbalance()?,
volume_weighted_price: trade_data.volume_weighted_price,
};
microstructure_metrics.push(microstructure_data);
order_flow_data.push(trade_data);
if tick % 1000 == 0 {
debug!("Tick {}: Spread = {:.4}, VPIN = {:.3}, Toxicity = {:.3}, Impact = {:.4}",
tick, spread_metrics.bid_ask_spread, vpin_score, toxicity_score, impact_metrics.temporary_impact);
}
}
}
}
// Calculate aggregate microstructure statistics
let avg_spread = microstructure_metrics.iter().map(|m| m.bid_ask_spread).sum::<f64>() / microstructure_metrics.len() as f64;
let avg_vpin = microstructure_metrics.iter().map(|m| m.vpin_score).sum::<f64>() / microstructure_metrics.len() as f64;
let avg_toxicity = microstructure_metrics.iter().map(|m| m.toxicity_score).sum::<f64>() / microstructure_metrics.len() as f64;
let avg_impact = microstructure_metrics.iter().map(|m| m.price_impact).sum::<f64>() / microstructure_metrics.len() as f64;
// Calculate market quality score (lower spreads and impacts = higher quality)
let market_quality_score = 1.0 / (1.0 + avg_spread + avg_impact);
Ok(ExampleMetrics {
score: Decimal::try_from(market_quality_score).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(avg_impact).unwrap_or(Decimal::ZERO)), // Market impact as "loss"
sharpe_ratio: Some(Decimal::try_from(1.0 - avg_vpin).unwrap_or(Decimal::ZERO)), // Inverse VPIN (lower = better)
max_drawdown: Some(Decimal::try_from(avg_toxicity).unwrap_or(Decimal::ZERO)), // Flow toxicity
})
}
*/
/// List all available examples
pub fn list_examples() -> Vec<ExampleType> {
vec![
ExampleType::BasicDQN,
ExampleType::RainbowDQN,
ExampleType::PriceTransformer,
ExampleType::RiskModels,
ExampleType::PortfolioOptimization,
ExampleType::Microstructure,
]
}
// Helper functions for examples
/// Simulate environment step for DQN training
fn simulate_environment_step(
state: &[f64],
action: crate::dqn::TradingAction,
) -> (Vec<f64>, f64, bool) {
// Simple environment simulation
let mut next_state = state.to_vec();
// Apply action effect
let action_value = match action {
crate::dqn::TradingAction::Buy => 1.0,
crate::dqn::TradingAction::Sell => -1.0,
crate::dqn::TradingAction::Hold => 0.0,
};
for i in 0..next_state.len() {
next_state[i] += action_value * random::<f64>() * 0.1;
// Ensure price fields (first 2 elements) are always positive for Price validation
if i < 2 {
next_state[i] = next_state[i].abs().max(0.01); // Minimum price of 0.01
}
}
// Calculate reward based on action and state change
let reward = match action {
crate::dqn::TradingAction::Buy | crate::dqn::TradingAction::Sell => {
random::<f64>() * 2.0 - 1.0
},
crate::dqn::TradingAction::Hold => random::<f64>() * 0.1,
};
// Episode ends randomly or based on conditions
let done = next_state.iter().any(|&x| x.abs() > 10.0) || random::<f64>() < 0.01;
(next_state, reward, done)
}
/// Calculate Sharpe ratio from loss values
fn calculate_sharpe_ratio(losses: &[f64]) -> f64 {
if losses.is_empty() {
return 0.0;
}
let mean_loss = losses.iter().sum::<f64>() / losses.len() as f64;
let std_loss = {
let variance =
losses.iter().map(|&x| (x - mean_loss).powi(2)).sum::<f64>() / losses.len() as f64;
variance.sqrt()
};
if std_loss == 0.0 {
0.0
} else {
-mean_loss / std_loss
} // Negative because we want lower loss
}
/// Calculate maximum drawdown from loss values
fn calculate_max_drawdown(losses: &[f64]) -> f64 {
if losses.is_empty() {
return 0.0;
}
let mut running_min = losses[0];
let mut max_drawdown: f64 = 0.0;
for &loss in losses {
running_min = running_min.min(loss);
max_drawdown = max_drawdown.max(loss - running_min);
}
max_drawdown
}
/// Calculate Sharpe ratio from reward values
fn calculate_sharpe_ratio_from_rewards(rewards: &[f64]) -> f64 {
if rewards.is_empty() {
return 0.0;
}
let mean_reward = rewards.iter().sum::<f64>() / rewards.len() as f64;
let std_reward = {
let variance = rewards
.iter()
.map(|&x| (x - mean_reward).powi(2))
.sum::<f64>()
/ rewards.len() as f64;
variance.sqrt()
};
if std_reward == 0.0 {
0.0
} else {
mean_reward / std_reward
}
}
/// Calculate maximum drawdown from reward values
fn calculate_max_drawdown_from_rewards(rewards: &[f64]) -> f64 {
if rewards.is_empty() {
return 0.0;
}
let mut running_max = rewards[0];
let mut max_drawdown: f64 = 0.0;
for &reward in rewards {
running_max = running_max.max(reward);
max_drawdown = max_drawdown.max(running_max - reward);
}
max_drawdown / running_max.abs().max(1.0) // Normalize by max value
}
// End of commented out sections
/// Run microstructure analysis example
async fn run_microstructure_example(_config: &ExampleConfig) -> Result<ExampleMetrics, MLError> {
// Placeholder implementation for microstructure analysis
// This would normally involve analyzing market microstructure patterns
info!("Running microstructure analysis example...");
// Return basic metrics for now
Ok(ExampleMetrics {
score: Decimal::try_from(0.75).unwrap_or(Decimal::ZERO),
loss: Some(Decimal::try_from(0.05).unwrap_or(Decimal::ZERO)),
sharpe_ratio: Some(Decimal::try_from(1.5).unwrap_or(Decimal::ZERO)),
max_drawdown: Some(Decimal::try_from(0.05).unwrap_or(Decimal::ZERO)),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_example_config_default() {
let config = ExampleConfig::default();
assert!(matches!(config.example_type, ExampleType::BasicDQN));
assert!(config.enable_safety);
}
#[tokio::test]
async fn test_run_basic_example() {
let config = ExampleConfig::default();
let result = run_example(config).await;
assert!(result.is_ok());
}
#[test]
fn test_list_examples() {
let examples = list_examples();
assert_eq!(examples.len(), 6);
}
}