From 20cbe7a7fa7a3d7f61811893195b082bc2fb2be5 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 21 Feb 2026 13:02:42 +0100 Subject: [PATCH] feat(ensemble): add ModelInferenceAdapter trait and types Co-Authored-By: Claude Opus 4.6 --- ml/src/ensemble/inference_adapter.rs | 96 ++++++++++++++++++++++++++++ ml/src/ensemble/mod.rs | 2 + 2 files changed, 98 insertions(+) create mode 100644 ml/src/ensemble/inference_adapter.rs diff --git a/ml/src/ensemble/inference_adapter.rs b/ml/src/ensemble/inference_adapter.rs new file mode 100644 index 000000000..2e0e9b59b --- /dev/null +++ b/ml/src/ensemble/inference_adapter.rs @@ -0,0 +1,96 @@ +//! Model inference adapter trait for ensemble prediction +//! +//! Defines the contract that each model adapter must implement +//! to participate in ensemble prediction. + +use crate::MLResult; + +/// Canonical feature vector for ensemble inference. +/// 51-dim layout matching DQN state_dim: +/// 0-9: price features (OHLCV, VWAP, spread, tick) +/// 10-19: technical indicators (RSI, MACD, Bollinger, ATR) +/// 20-29: order book (depth levels, imbalance) +/// 30-39: microstructure (trade flow, toxicity) +/// 40-50: position/risk (PnL, exposure, drawdown) +#[derive(Debug, Clone)] +pub struct FeatureVector { + pub values: Vec, + pub timestamp: i64, +} + +/// Prediction output from a single model adapter, normalized for ensemble aggregation. +#[derive(Debug, Clone)] +pub struct EnsemblePrediction { + pub model_name: String, + /// Directional signal: -1.0 (bearish) to 1.0 (bullish) + pub direction: f64, + /// Model confidence: 0.0 to 1.0 + pub confidence: f64, + /// Model-specific metadata + pub metadata: PredictionMeta, +} + +/// Optional model-specific metadata attached to predictions. +#[derive(Debug, Clone, Default)] +pub struct PredictionMeta { + /// Inference latency in microseconds + pub latency_us: u64, + /// Quantile forecasts (TFT only) + pub quantiles: Option>, + /// Attention weights (TFT only) + pub attention_weights: Option>, + /// Raw Q-values (DQN only) + pub q_values: Option>, +} + +/// Adapter trait for running inference on a loaded model. +/// +/// Each model type (DQN, PPO, TFT, Mamba2) implements this trait +/// to handle its own feature-to-tensor mapping and output normalization. +pub trait ModelInferenceAdapter: Send + Sync { + /// Human-readable model name for logging + fn model_name(&self) -> &str; + + /// Run inference on a canonical feature vector. + /// Returns a normalized ensemble prediction. + fn predict(&self, features: &FeatureVector) -> MLResult; + + /// Whether this adapter has a loaded model ready for inference + fn is_ready(&self) -> bool; +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_feature_vector_creation() { + let fv = FeatureVector { + values: vec![0.1; 51], + timestamp: 1700000000_000_000, + }; + assert_eq!(fv.values.len(), 51); + assert_eq!(fv.timestamp, 1700000000_000_000); + } + + #[test] + fn test_ensemble_prediction_direction_bounds() { + let pred = EnsemblePrediction { + model_name: "test".to_string(), + direction: 0.75, + confidence: 0.9, + metadata: PredictionMeta::default(), + }; + assert!(pred.direction >= -1.0 && pred.direction <= 1.0); + assert!(pred.confidence >= 0.0 && pred.confidence <= 1.0); + } + + #[test] + fn test_prediction_meta_default() { + let meta = PredictionMeta::default(); + assert_eq!(meta.latency_us, 0); + assert!(meta.quantiles.is_none()); + assert!(meta.attention_weights.is_none()); + assert!(meta.q_values.is_none()); + } +} diff --git a/ml/src/ensemble/mod.rs b/ml/src/ensemble/mod.rs index bbbd9f900..1861a8011 100644 --- a/ml/src/ensemble/mod.rs +++ b/ml/src/ensemble/mod.rs @@ -17,6 +17,7 @@ pub mod model; pub mod training_integration; // Training integration for ML service pub mod voting; pub mod weights; +pub mod inference_adapter; // Re-export key types that are used across ensemble modules pub use ab_testing::{ @@ -43,6 +44,7 @@ pub use metrics::{ CHECKPOINT_SWAP_LATENCY_MICROSECONDS, CHECKPOINT_VALIDATION_TOTAL, }; pub use training_integration::EnsembleTrainingIntegration; +pub use inference_adapter::{EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta}; /// Errors that can occur in ensemble operations #[derive(Error, Debug)]