//! Lightweight synchronous ensemble coordinator //! //! Wraps N [`ModelInferenceAdapter`] instances and aggregates their //! predictions via confidence-weighted voting. Models that are not //! ready or that fail inference are gracefully skipped with warnings. use std::collections::HashMap; use crate::ensemble::inference_adapter::{ EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, }; use crate::{MLError, MLResult}; /// Synchronous coordinator that aggregates predictions from multiple /// [`ModelInferenceAdapter`] instances using confidence-weighted voting. /// /// # Aggregation formula /// /// ```text /// weighted_direction = sum(direction_i * weight_i * confidence_i) /// / sum(weight_i * confidence_i) /// ``` /// /// where `weight_i` defaults to 1.0 and can be overridden via /// [`set_weight`](InferenceEnsemble::set_weight). #[allow(missing_debug_implementations)] pub struct InferenceEnsemble { adapters: Vec>, weights: HashMap, } impl InferenceEnsemble { /// Create a new ensemble from a vector of adapters. /// /// All adapters start with a default weight of 1.0. pub fn new(adapters: Vec>) -> Self { Self { adapters, weights: HashMap::new(), } } /// Set a custom weight for a model identified by name. /// /// Models without an explicit weight use 1.0. pub fn set_weight(&mut self, model_name: &str, weight: f64) { self.weights.insert(model_name.to_string(), weight); } /// Return the number of adapters whose `is_ready()` returns true. pub fn ready_count(&self) -> usize { self.adapters.iter().filter(|a| a.is_ready()).count() } /// Run inference across all ready adapters and aggregate via /// confidence-weighted voting. /// /// Returns [`MLError::InferenceError`] if no models are ready. pub fn predict(&self, features: &FeatureVector) -> MLResult { let ready_adapters: Vec<&Box> = self.adapters.iter().filter(|a| a.is_ready()).collect(); if ready_adapters.is_empty() { return Err(MLError::InferenceError( "No models are ready for inference".to_string(), )); } let mut weighted_direction_sum = 0.0_f64; let mut weight_confidence_sum = 0.0_f64; let mut confidence_sum = 0.0_f64; let mut successful_count = 0_usize; let mut model_names: Vec = Vec::new(); for adapter in &ready_adapters { let model_name = adapter.model_name().to_string(); match adapter.predict(features) { Ok(pred) => { // Circuit breaker: skip NaN/Inf predictions if !pred.direction.is_finite() || !pred.confidence.is_finite() { tracing::warn!( model = %model_name, direction = %pred.direction, confidence = %pred.confidence, "Model returned NaN/Inf prediction, skipping (circuit breaker)" ); continue; } // Clamp confidence to valid range let confidence = pred.confidence.clamp(0.0, 1.0); let w = self .weights .get(&model_name) .copied() .unwrap_or(1.0); let wc = w * confidence; weighted_direction_sum += pred.direction * wc; weight_confidence_sum += wc; confidence_sum += confidence; successful_count += 1; model_names.push(model_name); } Err(e) => { tracing::warn!( model = %model_name, error = %e, "Model prediction failed, skipping" ); } } } if successful_count == 0 { return Err(MLError::InferenceError( "All ready models failed during inference".to_string(), )); } let direction = if weight_confidence_sum.abs() < f64::EPSILON { 0.0 } else { weighted_direction_sum / weight_confidence_sum }; let avg_confidence = confidence_sum / successful_count as f64; let ensemble_name = format!("ENSEMBLE({})", model_names.join("+")); Ok(EnsemblePrediction { model_name: ensemble_name, direction, confidence: avg_confidence, metadata: PredictionMeta::default(), }) } } #[cfg(test)] mod tests { use super::*; /// A simple test adapter with configurable direction, confidence, and readiness. struct DummyAdapter { name: String, direction: f64, confidence: f64, ready: bool, } impl ModelInferenceAdapter for DummyAdapter { fn model_name(&self) -> &str { &self.name } fn predict(&self, _features: &FeatureVector) -> MLResult { Ok(EnsemblePrediction { model_name: self.name.clone(), direction: self.direction, confidence: self.confidence, metadata: PredictionMeta::default(), }) } fn is_ready(&self) -> bool { self.ready } } fn make_features() -> FeatureVector { FeatureVector { values: vec![0.0; 51], timestamp: 1_700_000_000_000_000, } } #[test] fn test_ensemble_equal_weight_aggregation() { // Model A: bullish (dir=1.0, conf=0.8) // Model B: bearish (dir=-1.0, conf=0.6) // weighted_direction = (1.0*1.0*0.8 + (-1.0)*1.0*0.6) / (1.0*0.8 + 1.0*0.6) // = (0.8 - 0.6) / 1.4 // = 0.2 / 1.4 // ~ 0.1429 let adapters: Vec> = vec![ Box::new(DummyAdapter { name: "A".to_string(), direction: 1.0, confidence: 0.8, ready: true, }), Box::new(DummyAdapter { name: "B".to_string(), direction: -1.0, confidence: 0.6, ready: true, }), ]; let ensemble = InferenceEnsemble::new(adapters); let features = make_features(); let pred = ensemble.predict(&features).expect("predict should succeed"); // Net direction should be positive (bullish model has higher confidence) assert!( pred.direction > 0.0, "direction should be positive, got {}", pred.direction ); // But not too strong since the bearish model partially cancels assert!( pred.direction < 0.5, "direction should be < 0.5, got {}", pred.direction ); // Confidence = average = (0.8 + 0.6) / 2 = 0.7 assert!( (pred.confidence - 0.7).abs() < 1e-9, "confidence should be 0.7, got {}", pred.confidence ); } #[test] fn test_ensemble_skips_unready_models() { let adapters: Vec> = vec![ Box::new(DummyAdapter { name: "Ready".to_string(), direction: 1.0, confidence: 0.9, ready: true, }), Box::new(DummyAdapter { name: "NotReady".to_string(), direction: -1.0, confidence: 0.9, ready: false, }), ]; let ensemble = InferenceEnsemble::new(adapters); assert_eq!(ensemble.ready_count(), 1); let features = make_features(); let pred = ensemble.predict(&features).expect("predict should succeed"); // Only the ready (bullish) model contributes assert!( pred.direction > 0.5, "direction should be > 0.5 with only bullish model, got {}", pred.direction ); // Model name should only include the ready model assert!( pred.model_name.contains("Ready"), "model_name should contain 'Ready', got {}", pred.model_name ); assert!( !pred.model_name.contains("NotReady"), "model_name should NOT contain 'NotReady', got {}", pred.model_name ); } #[test] fn test_ensemble_custom_weights() { // Model A: bullish, weight 0.75 // Model B: bearish, weight 0.25 // weighted_direction = (1.0*0.75*0.8 + (-1.0)*0.25*0.8) / (0.75*0.8 + 0.25*0.8) // = (0.6 - 0.2) / (0.6 + 0.2) // = 0.4 / 0.8 // = 0.5 let adapters: Vec> = vec![ Box::new(DummyAdapter { name: "Heavy".to_string(), direction: 1.0, confidence: 0.8, ready: true, }), Box::new(DummyAdapter { name: "Light".to_string(), direction: -1.0, confidence: 0.8, ready: true, }), ]; let mut ensemble = InferenceEnsemble::new(adapters); ensemble.set_weight("Heavy", 0.75); ensemble.set_weight("Light", 0.25); let features = make_features(); let pred = ensemble.predict(&features).expect("predict should succeed"); // The heavier-weighted bullish model should dominate assert!( pred.direction > 0.0, "direction should be positive (heavy bullish dominates), got {}", pred.direction ); // With equal confidences, the 3:1 weight ratio strongly favors bullish // Expected: 0.5, but any positive value > 0.3 confirms dominance assert!( pred.direction > 0.3, "direction should be > 0.3 showing heavy model dominance, got {}", pred.direction ); } #[test] fn test_ensemble_filters_nan_predictions() { // NaN model should be skipped; the valid model's prediction stands alone. let adapters: Vec> = vec![ Box::new(DummyAdapter { name: "Valid".to_string(), direction: 0.8, confidence: 0.7, ready: true, }), Box::new(DummyAdapter { name: "NaN_Model".to_string(), direction: f64::NAN, confidence: 0.9, ready: true, }), ]; let ensemble = InferenceEnsemble::new(adapters); let features = make_features(); let pred = ensemble.predict(&features).expect("predict should succeed"); // Only the Valid model should contribute assert!( pred.model_name.contains("Valid"), "model_name should contain 'Valid', got {}", pred.model_name ); assert!( !pred.model_name.contains("NaN_Model"), "model_name should NOT contain 'NaN_Model', got {}", pred.model_name ); // direction should come entirely from the Valid model assert!( (pred.direction - 0.8).abs() < 1e-9, "direction should be 0.8, got {}", pred.direction ); // confidence should be from the single valid model assert!( (pred.confidence - 0.7).abs() < 1e-9, "confidence should be 0.7, got {}", pred.confidence ); } #[test] fn test_ensemble_filters_extreme_confidence() { // Model with confidence 5.0 should be clamped to 1.0 let adapters: Vec> = vec![ Box::new(DummyAdapter { name: "Normal".to_string(), direction: 1.0, confidence: 0.6, ready: true, }), Box::new(DummyAdapter { name: "Extreme".to_string(), direction: 1.0, confidence: 5.0, ready: true, }), ]; let ensemble = InferenceEnsemble::new(adapters); let features = make_features(); let pred = ensemble.predict(&features).expect("predict should succeed"); // Both models should contribute (extreme confidence is clamped, not skipped) assert!( pred.model_name.contains("Normal"), "model_name should contain 'Normal', got {}", pred.model_name ); assert!( pred.model_name.contains("Extreme"), "model_name should contain 'Extreme', got {}", pred.model_name ); // Confidence should be clamped: avg of 0.6 and 1.0 = 0.8 assert!( (pred.confidence - 0.8).abs() < 1e-9, "confidence should be 0.8 (avg of 0.6 and clamped 1.0), got {}", pred.confidence ); } }