//! 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"); } }