diff --git a/backtesting/src/lib.rs b/backtesting/src/lib.rs index d5311a8bf..24fa72bba 100644 --- a/backtesting/src/lib.rs +++ b/backtesting/src/lib.rs @@ -69,6 +69,7 @@ use rust_decimal::Decimal; pub mod dbn_converter; pub mod dbn_replay; pub mod metrics; +pub mod model_loader; pub mod replay_engine; pub mod slippage; pub mod strategies; @@ -76,6 +77,7 @@ pub mod strategy_runner; pub mod strategy_tester; // Re-export types for public API +pub use model_loader::{load_models_for_backtest, BacktestMLModel, ModelSpec}; pub use metrics::{MetricsCalculator, PerformanceAnalytics}; pub use replay_engine::{MarketReplay, ReplayConfig}; pub use strategies::{DQNAction, DQNReplayStrategy, PositionState, TradingActionType}; diff --git a/backtesting/src/model_loader.rs b/backtesting/src/model_loader.rs new file mode 100644 index 000000000..561924ac1 --- /dev/null +++ b/backtesting/src/model_loader.rs @@ -0,0 +1,341 @@ +//! Model loading and MLModel wrapper for backtesting. +//! +//! Provides [`BacktestMLModel`], a bridge that wraps a +//! [`ModelInferenceAdapter`] (ensemble-level inference) and exposes +//! it through the global [`MLModel`] trait so it can be registered +//! in the [`ModelRegistry`]. + +use std::collections::HashMap; +use std::sync::Arc; + +use async_trait::async_trait; + +use ml::dqn::dqn::DQNConfig; +use ml::ensemble::adapters::dqn::DqnInferenceAdapter; +use ml::ensemble::adapters::ppo::PpoInferenceAdapter; +use ml::ensemble::inference_adapter::{FeatureVector, ModelInferenceAdapter}; +use ml::ppo::ppo::PPOConfig; +use ml::{Features, MLError, MLModel, MLResult, ModelMetadata, ModelPrediction, ModelType}; + +/// Wraps a [`ModelInferenceAdapter`] to satisfy the [`MLModel`] trait +/// expected by the global [`ModelRegistry`]. +/// +/// The conversion path is: +/// `Features` -> `FeatureVector` -> adapter `predict` -> `EnsemblePrediction` -> `ModelPrediction` +pub struct BacktestMLModel { + adapter: Box, + model_type: ModelType, +} + +impl std::fmt::Debug for BacktestMLModel { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BacktestMLModel") + .field("model_name", &self.adapter.model_name()) + .field("model_type", &self.model_type) + .field("is_ready", &self.adapter.is_ready()) + .finish() + } +} + +impl BacktestMLModel { + /// Create a new `BacktestMLModel` from an inference adapter and its type. + pub fn new(adapter: Box, model_type: ModelType) -> Self { + Self { + adapter, + model_type, + } + } + + /// Delegates to the underlying adapter's readiness check. + pub fn is_ready(&self) -> bool { + self.adapter.is_ready() + } +} + +#[async_trait] +impl MLModel for BacktestMLModel { + fn name(&self) -> &str { + self.adapter.model_name() + } + + fn model_type(&self) -> ModelType { + self.model_type + } + + async fn predict(&self, features: &Features) -> MLResult { + // Convert Features -> FeatureVector + let fv = FeatureVector { + values: features.values.clone(), + timestamp: features.timestamp as i64, + }; + + // Run inference through the adapter + let ensemble_pred = self.adapter.predict(&fv)?; + + // Build metadata map from the prediction metadata + let mut metadata = HashMap::new(); + metadata.insert( + "direction".to_string(), + serde_json::json!(ensemble_pred.direction), + ); + metadata.insert( + "latency_us".to_string(), + serde_json::json!(ensemble_pred.metadata.latency_us), + ); + if let Some(ref q_values) = ensemble_pred.metadata.q_values { + metadata.insert("q_values".to_string(), serde_json::json!(q_values)); + } + if let Some(ref quantiles) = ensemble_pred.metadata.quantiles { + metadata.insert("quantiles".to_string(), serde_json::json!(quantiles)); + } + if let Some(ref attn) = ensemble_pred.metadata.attention_weights { + metadata.insert("attention_weights".to_string(), serde_json::json!(attn)); + } + + Ok(ModelPrediction { + value: ensemble_pred.direction, + confidence: ensemble_pred.confidence, + metadata, + timestamp: features.timestamp, + model_id: ensemble_pred.model_name, + }) + } + + fn get_confidence(&self) -> f64 { + 0.8 + } + + fn is_ready(&self) -> bool { + self.adapter.is_ready() + } + + fn get_metadata(&self) -> ModelMetadata { + ModelMetadata::new( + self.model_type, + "1.0.0".to_string(), + 51, // canonical 51-dim feature vector + 0.0, + ) + } + + fn validate_features(&self, features: &Features) -> MLResult<()> { + if features.values.is_empty() { + return Err(MLError::ValidationError { + message: "Empty feature vector".to_string(), + }); + } + Ok(()) + } +} + +/// Specification for loading a model into the backtest registry. +/// +/// Each `ModelSpec` describes one model to instantiate: its type +/// (e.g. `"DQN"`, `"PPO"`), an optional checkpoint path, and the +/// ensemble weight used by the aggregation layer. +#[derive(Debug, Clone)] +pub struct ModelSpec { + /// Model type identifier (`"DQN"`, `"PPO"`). + pub model_type: String, + /// Path to a `.safetensors` checkpoint. Empty string = random weights. + pub checkpoint_path: String, + /// Ensemble weight (used downstream by the aggregation layer). + pub weight: f64, +} + +/// Load models described by `specs` and register them in the global +/// [`ModelRegistry`](ml::ModelRegistry). +/// +/// Call this **before** running a backtest so that the strategy layer +/// can look up models by name. +/// +/// # Errors +/// +/// Returns `MLError` if a model fails to initialise or register. +/// Unknown `model_type` values are logged and skipped. +pub async fn load_models_for_backtest(specs: &[ModelSpec]) -> Result<(), MLError> { + let registry = ml::get_global_registry(); + + for spec in specs { + let adapter: Box = match spec.model_type.as_str() { + "DQN" => { + let config = DQNConfig { + state_dim: 51, + num_actions: 45, + hidden_dims: vec![128, 128], + ..Default::default() + }; + if spec.checkpoint_path.is_empty() { + Box::new(DqnInferenceAdapter::new(config)?) + } else { + Box::new(DqnInferenceAdapter::from_checkpoint( + config, + &spec.checkpoint_path, + )?) + } + } + "PPO" => { + let config = PPOConfig { + state_dim: 51, + num_actions: 45, + policy_hidden_dims: vec![128, 64], + value_hidden_dims: vec![128, 64], + ..Default::default() + }; + Box::new(PpoInferenceAdapter::new(config)?) + } + other => { + tracing::warn!("Unknown model type: {}, skipping", other); + continue; + } + }; + + let model_type = match spec.model_type.as_str() { + "DQN" => ModelType::DQN, + "PPO" => ModelType::PPO, + _ => { + // Already skipped above via `continue`, but guard against + // future additions. + tracing::warn!("Unmapped model type: {}, skipping registration", spec.model_type); + continue; + } + }; + + let ml_model = BacktestMLModel::new(adapter, model_type); + registry.register(Arc::new(ml_model)).await?; + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use ml::dqn::dqn::DQNConfig; + use ml::ensemble::adapters::dqn::DqnInferenceAdapter; + + fn test_dqn_config() -> DQNConfig { + DQNConfig { + state_dim: 51, + num_actions: 45, + hidden_dims: vec![64, 64], + ..Default::default() + } + } + + #[tokio::test] + async fn test_backtest_ml_model_wraps_dqn() { + let adapter = DqnInferenceAdapter::new(test_dqn_config()) + .expect("DqnInferenceAdapter creation should succeed in test"); + let model = BacktestMLModel::new(Box::new(adapter), ModelType::DQN); + + // Build a 51-dim feature set + let features = Features::new(vec![0.1; 51], (0..51).map(|i| format!("f{i}")).collect()); + + let prediction = model + .predict(&features) + .await + .expect("predict should succeed in test"); + + // Direction must be in [-1, 1] + assert!( + prediction.value >= -1.0 && prediction.value <= 1.0, + "direction {} out of [-1, 1]", + prediction.value, + ); + // Confidence must be in [0, 1] + assert!( + prediction.confidence >= 0.0 && prediction.confidence <= 1.0, + "confidence {} out of [0, 1]", + prediction.confidence, + ); + // Model id should be "DQN" + assert_eq!(prediction.model_id, "DQN"); + } + + #[tokio::test] + async fn test_backtest_ml_model_name_and_type() { + let adapter = DqnInferenceAdapter::new(test_dqn_config()) + .expect("DqnInferenceAdapter creation should succeed in test"); + let model = BacktestMLModel::new(Box::new(adapter), ModelType::DQN); + + assert_eq!(model.name(), "DQN"); + assert_eq!(model.model_type(), ModelType::DQN); + assert!(model.is_ready()); + } + + #[tokio::test] + async fn test_backtest_ml_model_validate_empty_features() { + let adapter = DqnInferenceAdapter::new(test_dqn_config()) + .expect("DqnInferenceAdapter creation should succeed in test"); + let model = BacktestMLModel::new(Box::new(adapter), ModelType::DQN); + + let empty_features = Features::new(vec![], vec![]); + let result = model.validate_features(&empty_features); + assert!(result.is_err(), "should reject empty features"); + } + + #[test] + fn test_model_spec_creation() { + let spec = ModelSpec { + model_type: "DQN".to_string(), + checkpoint_path: String::new(), + weight: 0.6, + }; + assert_eq!(spec.model_type, "DQN"); + assert!(spec.checkpoint_path.is_empty()); + assert!((spec.weight - 0.6).abs() < f64::EPSILON); + } + + #[tokio::test] + async fn test_load_dqn_into_registry() { + let specs = vec![ModelSpec { + model_type: "DQN".to_string(), + checkpoint_path: String::new(), + weight: 1.0, + }]; + + load_models_for_backtest(&specs) + .await + .expect("load_models_for_backtest should succeed in test"); + + let registry = ml::get_global_registry(); + let model = registry.get("DQN").await; + assert!(model.is_some(), "DQN should be registered in the global registry"); + + let model = model.expect("already checked Some above"); + assert_eq!(model.name(), "DQN"); + assert_eq!(model.model_type(), ModelType::DQN); + assert!(model.is_ready()); + } + + #[tokio::test] + async fn test_load_ppo_into_registry() { + let specs = vec![ModelSpec { + model_type: "PPO".to_string(), + checkpoint_path: String::new(), + weight: 0.5, + }]; + + load_models_for_backtest(&specs) + .await + .expect("load_models_for_backtest should succeed for PPO in test"); + + let registry = ml::get_global_registry(); + let model = registry.get("PPO").await; + assert!(model.is_some(), "PPO should be registered in the global registry"); + } + + #[tokio::test] + async fn test_load_unknown_model_type_skipped() { + let specs = vec![ModelSpec { + model_type: "UNKNOWN_MODEL".to_string(), + checkpoint_path: String::new(), + weight: 1.0, + }]; + + // Should succeed (unknown types are skipped, not errored) + let result = load_models_for_backtest(&specs).await; + assert!(result.is_ok(), "unknown model types should be skipped without error"); + } +}