From c300fa0551057a5357bf711d45fcd888dff12959 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 21 Feb 2026 13:35:28 +0100 Subject: [PATCH] feat(ensemble): integrate real inference adapters into coordinator Replace mock predictions with real model inference when adapters are registered. Falls back to mock predictions for models without adapters. Co-Authored-By: Claude Opus 4.6 --- ml/src/ensemble/coordinator.rs | 124 ++++++++++++++++++++++++++++++++- 1 file changed, 123 insertions(+), 1 deletion(-) diff --git a/ml/src/ensemble/coordinator.rs b/ml/src/ensemble/coordinator.rs index b853bf730..2511ecbc0 100644 --- a/ml/src/ensemble/coordinator.rs +++ b/ml/src/ensemble/coordinator.rs @@ -4,6 +4,7 @@ //! from multiple ML models (DQN, PPO, TFT, MAMBA-2, Liquid, TLOB) for production trading decisions. //! Supports dynamic weighting based on performance and model diversity metrics. +use crate::ensemble::inference_adapter::{FeatureVector, ModelInferenceAdapter}; use crate::ensemble::{EnsembleDecision, ModelVote, ModelWeight, TradingAction}; use crate::{Features, MLError, MLResult, ModelPrediction}; use std::collections::HashMap; @@ -18,7 +19,6 @@ const MIN_WEIGHT_THRESHOLD: f64 = 0.05; const MAX_WEIGHT_THRESHOLD: f64 = 0.40; /// Ensemble coordinator for aggregating model predictions -#[derive(Debug)] pub struct EnsembleCoordinator { /// Active model registry (dual-buffer for hot-swapping) active_models: Arc>, @@ -31,6 +31,18 @@ pub struct EnsembleCoordinator { /// Configuration for ensemble behavior config: EnsembleConfig, + + /// Real model inference adapters (tried before mock predictions) + adapters: Vec>, +} + +impl std::fmt::Debug for EnsembleCoordinator { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("EnsembleCoordinator") + .field("config", &self.config) + .field("adapter_count", &self.adapters.len()) + .finish() + } } /// Configuration for ensemble coordinator @@ -64,9 +76,18 @@ impl EnsembleCoordinator { aggregator: Arc::new(SignalAggregator::new()), model_weights: Arc::new(RwLock::new(HashMap::new())), config, + adapters: Vec::new(), } } + /// Add a real model inference adapter to the ensemble. + /// When an adapter's model_name matches a registered model, its predictions + /// are used instead of mock predictions. + pub fn add_adapter(&mut self, adapter: Box) { + info!("Added inference adapter: {}", adapter.model_name()); + self.adapters.push(adapter); + } + /// Register a model in the ensemble pub async fn register_model(&self, model_id: String, weight: f64) -> MLResult<()> { let model_weight = ModelWeight::new(model_id.clone(), weight); @@ -127,6 +148,35 @@ impl EnsembleCoordinator { let mut predictions = Vec::new(); for (model_id, checkpoint_opt) in model_info { + // Try real adapter first + if let Some(adapter) = self + .adapters + .iter() + .find(|a| a.model_name() == model_id && a.is_ready()) + { + let fv = FeatureVector { + values: features.values.clone(), + timestamp: features.timestamp as i64, + }; + match adapter.predict(&fv) { + Ok(ensemble_pred) => { + let prediction = ModelPrediction::new( + model_id.clone(), + ensemble_pred.direction, + ensemble_pred.confidence, + ); + predictions.push(prediction); + continue; + } + Err(e) => { + warn!( + "Adapter {} inference failed, falling back to mock: {}", + model_id, e + ); + } + } + } + // Try to get real model from registry if let Some(checkpoint_path) = checkpoint_opt { debug!( @@ -661,4 +711,76 @@ mod tests { assert!(registry.active.contains_key("DQN")); } + + #[tokio::test] + async fn test_ensemble_with_real_adapter() { + use crate::ensemble::inference_adapter::{ + EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, + }; + + struct TestAdapter; + impl ModelInferenceAdapter for TestAdapter { + fn model_name(&self) -> &str { + "DQN" + } + fn predict( + &self, + _features: &FeatureVector, + ) -> crate::MLResult { + Ok(EnsemblePrediction { + model_name: "DQN".to_string(), + direction: 0.6, + confidence: 0.85, + metadata: PredictionMeta::default(), + }) + } + fn is_ready(&self) -> bool { + true + } + } + // Safety: TestAdapter has no mutable state, safe to share across threads + unsafe impl Send for TestAdapter {} + unsafe impl Sync for TestAdapter {} + + let mut coordinator = EnsembleCoordinator::new(); + coordinator.add_adapter(Box::new(TestAdapter)); + coordinator + .register_model("DQN".to_string(), 1.0) + .await + .unwrap(); + + let features = Features::new( + vec![0.5; 10], + vec!["f1", "f2", "f3", "f4", "f5", "f6", "f7", "f8", "f9", "f10"] + .into_iter() + .map(String::from) + .collect(), + ); + let decision = coordinator.predict(&features).await.unwrap(); + assert!(decision.confidence > 0.0); + assert_eq!(decision.model_count(), 1); + } + + #[tokio::test] + async fn test_ensemble_graceful_no_adapters() { + let coordinator = EnsembleCoordinator::new(); + coordinator + .register_model("DQN".to_string(), 0.5) + .await + .unwrap(); + coordinator + .register_model("PPO".to_string(), 0.5) + .await + .unwrap(); + + let features = Features::new( + vec![0.5; 5], + vec!["f1", "f2", "f3", "f4", "f5"] + .into_iter() + .map(String::from) + .collect(), + ); + let decision = coordinator.predict(&features).await.unwrap(); + assert!(decision.confidence >= 0.0); + } }