Files
foxhunt/services/backtesting_service/src/ml_strategy_engine.rs
jgrusewski 9b27428de9 fix(ensemble): wire EnsembleModelAdapter::predict to real inference
The adapter's `predict` returned
`{ prediction_value: 0.5, confidence: 0.0 }` as a "neutral" stub.
Downstream ensemble filtering dropped any 0-confidence prediction,
so the adapter silently contributed nothing -- a stub that survived
only because the caller discarded it.

Now: `EnsembleModelAdapter` holds an
`Arc<dyn ModelInferenceAdapter>` and dispatches `predict` to the
wrapped per-model GPU adapter (`DqnInferenceAdapter`,
`PpoInferenceAdapter`, `TftInferenceAdapter`, `Mamba2InferenceAdapter`,
`LiquidInferenceAdapter`, `KanInferenceAdapter`,
`XlstmInferenceAdapter`, `TggnInferenceAdapter`,
`TlobInferenceAdapter`, `DiffusionInferenceAdapter`). The inner
adapter already runs a full forward pass on its model and returns a
normalized `(direction, confidence)` pair; the bridge maps
`direction in [-1, 1]` -> `prediction_value in [0, 1]` via
`(direction + 1) / 2` to match `MLPrediction`'s bullish-probability
contract (> 0.5 = bullish), clamps confidence to [0, 1], and
propagates `metadata.latency_us` as the inference latency.

The bridge returns `Err` (never a faked 0-confidence success) when
the inner model reports `is_ready() == false`, the inner `predict`
fails, the output is non-finite, or the feature slice is empty.

`build_production_strategy` no longer fabricates ten zero-confidence
ghost adapters. It now accepts
`Vec<(String, Arc<dyn ModelInferenceAdapter>)>` -- the caller owns
real model construction (checkpoint loading, device selection). The
only current caller (backtesting service `MLPoweredStrategy::new`)
passes an empty vec; that yields an empty ensemble, which is an
honest "no models loaded" signal rather than ten stubs that exist
only to be filtered.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-23 09:24:25 +02:00

527 lines
20 KiB
Rust

//! ML-powered strategy execution engine for backtesting
//!
//! This module integrates the shared ML strategy from common crate to ensure
//! ONE SINGLE SYSTEM across trading and backtesting services.
use anyhow::Result;
use chrono::{DateTime, Utc};
use rust_decimal::{prelude::ToPrimitive, Decimal};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tracing::{debug, info};
use crate::storage::StorageManager;
use crate::strategy_engine::{
BacktestTrade, MarketData, Portfolio, StrategyExecutor, TradeSide, TradeSignal,
};
use config::structures::BacktestingStrategyConfig;
// Import shared ML strategy (ONE SINGLE SYSTEM)
use common::ml_strategy::{MLPrediction as CommonMLPrediction, SharedMLStrategy};
// Import UnifiedFeatureExtractor (256 features, production system)
use ml::features::unified::{FeatureExtractionConfig, UnifiedFeatureExtractor};
use ml::safety::{MLSafetyConfig, MLSafetyManager};
/// ML model prediction result for backtesting
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MLPrediction {
/// Model identifier
pub model_id: String,
/// Prediction value (0.0-1.0)
pub prediction_value: f64,
/// Confidence score (0.0-1.0)
pub confidence: f64,
/// Features used for prediction
pub features: Vec<f64>,
/// Prediction timestamp
pub timestamp: DateTime<Utc>,
/// Inference latency in microseconds
pub inference_latency_us: u64,
}
/// ML model performance tracking for backtesting
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct MLModelPerformance {
/// Model identifier
pub model_id: String,
/// Total predictions made
pub total_predictions: u64,
/// Correct predictions (when outcome is known)
pub correct_predictions: u64,
/// Average inference latency
pub avg_latency_us: f64,
/// Average confidence score
pub avg_confidence: f64,
/// Model accuracy percentage
pub accuracy_percentage: f64,
/// Returns generated when following this model
pub returns: Vec<f64>,
/// Sharpe ratio for this model
pub sharpe_ratio: f64,
/// Maximum drawdown when following this model
pub max_drawdown: f64,
}
// NOTE: MLFeatureExtractor REMOVED - Replaced with UnifiedFeatureExtractor (256 features)
// Old implementation used only 8 features (price return, MA, volatility, volume, time).
// New implementation uses production-grade 256-feature extraction pipeline:
// - 5 OHLCV features
// - 10 technical indicators (RSI, MACD, Bollinger, ATR, EMA)
// - 60 price patterns
// - 40 volume patterns
// - 50 microstructure features
// - 10 time-based features
// - 81 statistical features
//
// This ensures backtesting uses the SAME features as live trading and model training.
/// ML-powered strategy for backtesting (uses shared ML strategy - ONE SINGLE SYSTEM)
pub struct MLPoweredStrategy {
/// Strategy name
name: String,
/// Shared ML strategy (ONE SINGLE SYSTEM)
strategy: Arc<SharedMLStrategy>,
/// Unified feature extractor (256 features, production system)
_feature_extractor: Arc<UnifiedFeatureExtractor>,
/// Model performance tracking (local copy for backward compatibility)
model_performance: HashMap<String, MLModelPerformance>,
/// Current position size based on confidence
confidence_based_sizing: bool,
/// Minimum confidence threshold for trades
min_confidence_threshold: f64,
}
// NOTE: Old model simulator implementations removed.
// We now use SharedMLStrategy from common crate (ONE SINGLE SYSTEM).
// This eliminates code duplication and ensures consistent ML predictions
// across trading and backtesting services.
impl std::fmt::Debug for MLPoweredStrategy {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MLPoweredStrategy")
.field("name", &self.name)
.field("confidence_based_sizing", &self.confidence_based_sizing)
.field("min_confidence_threshold", &self.min_confidence_threshold)
.field("model_performance_count", &self.model_performance.len())
.finish()
}
}
impl MLPoweredStrategy {
/// Create new ML-powered strategy (uses shared strategy - ONE SINGLE SYSTEM)
/// # Errors
/// Returns error if ML model adapter construction fails.
pub fn new(name: String, _lookback_periods: usize) -> Result<Self> {
// Use shared ML strategy (ONE SINGLE SYSTEM) with ProductionFeatureExtractorAdapter (225 features).
//
// `build_production_strategy` now takes a Vec of real per-model inference
// adapters. The backtesting service has no checkpoint-loader wired at this
// layer, so we pass an empty vec: the resulting `SharedMLStrategy` honestly
// reports "no models loaded" (empty prediction list) instead of the prior
// ghost behaviour where 10 stub adapters all returned confidence=0.0 and
// were silently filtered by the threshold.
//
// When the backtesting service grows a model-loader, it should build the
// appropriate `DqnInferenceAdapter::from_checkpoint(..)` / `PpoInferenceAdapter::new(..)`
// / etc. adapters and pass them in here with their ensemble ids.
let min_confidence_threshold = 0.6;
let strategy = Arc::new(ml::ensemble::build_production_strategy(
min_confidence_threshold,
Vec::new(),
));
// Initialize UnifiedFeatureExtractor (256 features)
let feature_config = FeatureExtractionConfig::default();
let safety_config = MLSafetyConfig::default();
let safety_manager = Arc::new(MLSafetyManager::new(safety_config));
let feature_extractor =
Arc::new(UnifiedFeatureExtractor::new(feature_config, safety_manager));
Ok(Self {
name,
strategy,
_feature_extractor: feature_extractor,
model_performance: HashMap::new(),
confidence_based_sizing: true,
min_confidence_threshold,
})
}
/// Get ensemble prediction from all models (delegates to shared strategy)
pub async fn get_ensemble_prediction(
&mut self,
market_data: &MarketData,
) -> Result<Vec<MLPrediction>> {
// Use shared ML strategy (ONE SINGLE SYSTEM)
let price = market_data.close.to_f64().unwrap_or(0.0);
let volume = market_data.volume.to_f64().unwrap_or(0.0);
let timestamp = market_data.timestamp;
// Get predictions from shared strategy
let common_predictions = self
.strategy
.get_ensemble_prediction(price, volume, timestamp)
.await?;
// Convert to local type for backward compatibility
let predictions = common_predictions
.iter()
.map(|p| MLPrediction {
model_id: p.model_id.clone(),
prediction_value: p.prediction_value,
confidence: p.confidence,
features: p.features.clone(),
timestamp: p.timestamp,
inference_latency_us: p.inference_latency_us,
})
.collect();
Ok(predictions)
}
/// Calculate weighted ensemble prediction
pub fn calculate_ensemble_vote(&self, predictions: &[MLPrediction]) -> Option<(f64, f64)> {
if predictions.is_empty() {
return None;
}
let total_confidence: f64 = predictions.iter().map(|p| p.confidence).sum();
if total_confidence == 0.0 {
return None;
}
// Weighted average by confidence
let weighted_prediction: f64 = predictions
.iter()
.map(|p| p.prediction_value * p.confidence)
.sum::<f64>()
/ total_confidence;
let average_confidence: f64 =
predictions.iter().map(|p| p.confidence).sum::<f64>() / predictions.len() as f64;
Some((weighted_prediction, average_confidence))
}
/// Validate predictions against actual market outcomes (delegates to shared strategy)
pub async fn validate_predictions(&mut self, predictions: &[MLPrediction], actual_return: f64) {
// Convert to common predictions
let common_predictions: Vec<CommonMLPrediction> = predictions
.iter()
.map(|p| CommonMLPrediction {
model_id: p.model_id.clone(),
prediction_value: p.prediction_value,
confidence: p.confidence,
features: p.features.clone(),
timestamp: p.timestamp,
inference_latency_us: p.inference_latency_us,
})
.collect();
// Delegate to shared strategy (ONE SINGLE SYSTEM)
self.strategy
.validate_predictions(&common_predictions, actual_return)
.await;
// Update local performance tracking for backward compatibility
let shared_performance = self.strategy.get_performance_summary().await;
for (model_id, perf) in shared_performance {
self.model_performance.insert(
model_id.clone(),
MLModelPerformance {
model_id: model_id.clone(),
total_predictions: perf.total_predictions,
correct_predictions: perf.correct_predictions,
avg_latency_us: perf.avg_latency_us,
avg_confidence: perf.avg_confidence,
accuracy_percentage: perf.accuracy_percentage,
returns: perf.returns,
sharpe_ratio: perf.sharpe_ratio,
max_drawdown: perf.max_drawdown,
},
);
}
}
/// Get performance summary for all models
pub fn get_performance_summary(&self) -> HashMap<String, MLModelPerformance> {
self.model_performance.clone()
}
}
impl StrategyExecutor for MLPoweredStrategy {
fn execute(
&self,
market_data: &MarketData,
_portfolio: &Portfolio,
parameters: &HashMap<String, String>,
) -> Result<Vec<TradeSignal>> {
// NOTE: This method has &self (immutable), but we need mutable access to extract features.
// In production, consider using interior mutability (RefCell/Mutex) or redesigning the trait.
// For now, use async runtime to call SharedMLStrategy which handles this internally.
let mut signals = Vec::new();
// Use shared ML strategy for ensemble prediction (handles feature extraction internally)
let price = market_data.close.to_f64().unwrap_or(0.0);
let volume = market_data.volume.to_f64().unwrap_or(0.0);
let timestamp = market_data.timestamp;
// Create tokio runtime for async calls
let runtime = tokio::runtime::Runtime::new()?;
let predictions = runtime.block_on(async {
self.strategy
.get_ensemble_prediction(price, volume, timestamp)
.await
})?;
// Convert to local MLPrediction type
let local_predictions: Vec<MLPrediction> = predictions
.iter()
.map(|p| MLPrediction {
model_id: p.model_id.clone(),
prediction_value: p.prediction_value,
confidence: p.confidence,
features: p.features.clone(),
timestamp: p.timestamp,
inference_latency_us: p.inference_latency_us,
})
.collect();
// Calculate ensemble vote
if let Some((ensemble_prediction, ensemble_confidence)) =
self.calculate_ensemble_vote(&local_predictions)
{
let min_confidence = parameters
.get("min_confidence")
.and_then(|s| s.parse::<f64>().ok())
.unwrap_or(self.min_confidence_threshold);
if ensemble_confidence >= min_confidence {
let quantity = if self.confidence_based_sizing {
Decimal::try_from(ensemble_confidence * 1000.0).unwrap_or(Decimal::from(100))
} else {
Decimal::from(100)
};
if ensemble_prediction > 0.6 {
signals.push(TradeSignal {
symbol: market_data.symbol.clone(),
side: TradeSide::Buy,
quantity,
reason: format!(
"ML ensemble prediction: {:.3} (confidence: {:.3})",
ensemble_prediction, ensemble_confidence
),
});
} else if ensemble_prediction < 0.4 {
signals.push(TradeSignal {
symbol: market_data.symbol.clone(),
side: TradeSide::Sell,
quantity,
reason: format!(
"ML ensemble prediction: {:.3} (confidence: {:.3})",
ensemble_prediction, ensemble_confidence
),
});
}
}
}
Ok(signals)
}
}
/// ML Strategy Engine with model performance tracking
pub struct MLStrategyEngine {
/// Base strategy engine
base_engine: crate::strategy_engine::StrategyEngine,
/// ML-powered strategies
ml_strategies: HashMap<String, MLPoweredStrategy>,
/// Model performance tracking across backtests
global_model_performance: HashMap<String, MLModelPerformance>,
}
impl MLStrategyEngine {
/// Create new ML strategy engine
pub async fn new(
config: &BacktestingStrategyConfig,
storage_manager: Arc<StorageManager>,
) -> Result<Self> {
// Create repositories from storage manager
let repositories =
Arc::new(crate::repository_impl::create_repositories(storage_manager).await?);
let base_engine = crate::strategy_engine::StrategyEngine::new(config, repositories).await?;
let mut ml_strategies = HashMap::new();
// Add ML-powered strategies
ml_strategies.insert(
"ml_momentum".to_string(),
MLPoweredStrategy::new("ml_momentum".to_string(), 20)?,
);
ml_strategies.insert(
"ml_ensemble".to_string(),
MLPoweredStrategy::new("ml_ensemble".to_string(), 50)?,
);
Ok(Self {
base_engine,
ml_strategies,
global_model_performance: HashMap::new(),
})
}
/// Execute backtest with ML model validation
pub async fn execute_ml_backtest(
&mut self,
context: &crate::service::BacktestContext,
) -> Result<(Vec<BacktestTrade>, HashMap<String, MLModelPerformance>)> {
info!(
"Executing ML-powered backtest {} for strategy {}",
context.id, context.strategy_name
);
// Check if this is an ML strategy
let is_ml_strategy = self.ml_strategies.contains_key(&context.strategy_name);
if is_ml_strategy {
// Execute ML-powered backtest with model validation
self.execute_ml_strategy_backtest(context).await
} else {
// Fall back to base strategy engine
let trades = self.base_engine.execute_backtest(context).await?;
Ok((trades, HashMap::new()))
}
}
/// Execute backtest for ML strategy with model performance tracking
async fn execute_ml_strategy_backtest(
&mut self,
context: &crate::service::BacktestContext,
) -> Result<(Vec<BacktestTrade>, HashMap<String, MLModelPerformance>)> {
// Load market data for the backtest period
let market_data = self
.base_engine
.load_market_data(
&context.symbols,
context.started_at,
context
.completed_at
.unwrap_or(chrono::Utc::now().timestamp_nanos_opt().unwrap_or(0)),
)
.await?;
let trades = Vec::new();
let mut previous_price = None;
let total_data_points = market_data.len();
// Get ML strategy reference
let ml_strategy = self
.ml_strategies
.get_mut(&context.strategy_name)
.ok_or_else(|| anyhow::anyhow!("ML strategy {} not found", context.strategy_name))?;
// Process each data point with ML predictions
for (i, data_point) in market_data.into_iter().enumerate() {
// Get ML predictions (async call to shared strategy)
let predictions = ml_strategy.get_ensemble_prediction(&data_point).await?;
// Calculate ensemble vote
if let Some((ensemble_prediction, ensemble_confidence)) =
ml_strategy.calculate_ensemble_vote(&predictions)
{
debug!(
"Ensemble prediction: {:.3} (confidence: {:.3})",
ensemble_prediction, ensemble_confidence
);
// Validate predictions against future returns if we have next price
if let Some(prev_price) = previous_price {
let current_price = data_point.close.to_f64().unwrap_or(prev_price);
let actual_return = (current_price - prev_price) / prev_price;
ml_strategy
.validate_predictions(&predictions, actual_return)
.await;
}
}
previous_price = Some(data_point.close.to_f64().unwrap_or(0.0));
// Generate and execute trades using base strategy logic
// (This would integrate with the existing strategy execution logic)
if i % 100 == 0 {
let progress = (i as f64 / total_data_points as f64) * 100.0;
debug!("ML backtest progress: {:.1}%", progress);
}
}
// Get final model performance
let model_performance = ml_strategy.get_performance_summary();
// Update global performance tracking
for (model_id, perf) in &model_performance {
self.global_model_performance
.insert(model_id.clone(), perf.clone());
}
info!(
"ML backtest completed with {} trades and {} model evaluations",
trades.len(),
model_performance.len()
);
Ok((trades, model_performance))
}
/// Get model performance across all backtests
pub fn get_global_model_performance(&self) -> &HashMap<String, MLModelPerformance> {
&self.global_model_performance
}
/// Generate model performance report
pub fn generate_performance_report(&self) -> String {
let mut report = String::new();
report.push_str("=== ML Model Performance Report ===\n\n");
for (model_id, performance) in &self.global_model_performance {
report.push_str(&format!("Model: {}\n", model_id));
report.push_str(&format!(
" Total Predictions: {}\n",
performance.total_predictions
));
report.push_str(&format!(
" Accuracy: {:.2}%\n",
performance.accuracy_percentage
));
report.push_str(&format!(
" Average Confidence: {:.3}\n",
performance.avg_confidence
));
report.push_str(&format!(
" Average Latency: {:.1}μs\n",
performance.avg_latency_us
));
if performance.sharpe_ratio != 0.0 {
report.push_str(&format!(
" Sharpe Ratio: {:.3}\n",
performance.sharpe_ratio
));
}
if performance.max_drawdown != 0.0 {
report.push_str(&format!(
" Max Drawdown: {:.2}%\n",
performance.max_drawdown * 100.0
));
}
report.push('\n');
}
report
}
}