From 78f5ead601c3cfae6f4f5ef99e6e7772b77f3bea Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 2 Mar 2026 17:46:09 +0100 Subject: [PATCH] feat(ml): implement predict_raw() for 5 scalar-output ensemble adapters Override predict_raw() in TGGN, TLOB, KAN, xLSTM, Diffusion adapters to return raw GPU tensors. Enables GPU-side ensemble aggregation instead of per-model CPU extraction. Co-Authored-By: Claude Opus 4.6 --- crates/ml/src/ensemble/adapters/diffusion.rs | 36 +++++++++- crates/ml/src/ensemble/adapters/kan.rs | 33 ++++++++- crates/ml/src/ensemble/adapters/tggn.rs | 33 ++++++++- crates/ml/src/ensemble/adapters/tlob.rs | 69 ++++++++++++++++++- crates/ml/src/ensemble/adapters/xlstm.rs | 72 +++++++++++++++++++- 5 files changed, 238 insertions(+), 5 deletions(-) diff --git a/crates/ml/src/ensemble/adapters/diffusion.rs b/crates/ml/src/ensemble/adapters/diffusion.rs index 32b54973b..2129b22b1 100644 --- a/crates/ml/src/ensemble/adapters/diffusion.rs +++ b/crates/ml/src/ensemble/adapters/diffusion.rs @@ -12,7 +12,7 @@ use candle_nn::{VarBuilder, VarMap}; use crate::diffusion::config::DiffusionConfig; use crate::diffusion::denoiser::Denoiser; use crate::ensemble::inference_adapter::{ - EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, + EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; use crate::gpu::DeviceConfig; use crate::{MLError, MLResult}; @@ -155,6 +155,40 @@ impl ModelInferenceAdapter for DiffusionInferenceAdapter { }) } + fn predict_raw(&self, features: &FeatureVector) -> MLResult { + let padded = self.pad_features(&features.values); + let input = Tensor::from_vec(padded, (1, self.data_dim), &self.device) + .map_err(|e| MLError::ModelError(format!("Diffusion input tensor: {e}")))?; + + // Timestep t=1 (minimal noise level for feature processing) + let t = Tensor::from_vec(vec![1_u32], (1,), &self.device) + .map_err(|e| MLError::ModelError(format!("Diffusion timestep tensor: {e}")))?; + + let model = self + .model + .lock() + .map_err(|e| MLError::LockError(format!("Diffusion lock poisoned: {e}")))?; + let output = model.forward(&input, &t)?; + + // Mean of denoiser output -> scalar tensor (stays on GPU) + let mean_tensor = output + .mean_all() + .map_err(|e| MLError::ModelError(format!("Diffusion mean: {e}")))?; + + // One sync for confidence (acceptable) + let raw_f32: f32 = mean_tensor + .to_scalar() + .map_err(|e| MLError::ModelError(format!("Diffusion confidence calc: {e}")))?; + let prob = 1.0 / (1.0 + (-raw_f32 as f64).exp()); + let confidence = ((prob - 0.5).abs() * 2.0).clamp(0.0, 1.0); + + Ok(RawPrediction { + direction_scalar: 0.0, + confidence, + tensor: Some(mean_tensor), + }) + } + fn is_ready(&self) -> bool { self.model.lock().is_ok() } diff --git a/crates/ml/src/ensemble/adapters/kan.rs b/crates/ml/src/ensemble/adapters/kan.rs index 0fb0baead..7a57a0f7c 100644 --- a/crates/ml/src/ensemble/adapters/kan.rs +++ b/crates/ml/src/ensemble/adapters/kan.rs @@ -9,7 +9,7 @@ use candle_core::{DType, Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; use crate::ensemble::inference_adapter::{ - EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, + EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; use crate::gpu::DeviceConfig; use crate::kan::config::KANConfig; @@ -134,6 +134,37 @@ impl ModelInferenceAdapter for KanInferenceAdapter { }) } + fn predict_raw(&self, features: &FeatureVector) -> MLResult { + let padded = self.pad_features(&features.values); + let input = Tensor::from_vec(padded, (1, self.input_dim), &self.device) + .map_err(|e| MLError::ModelError(format!("KAN input tensor: {e}")))?; + + let model = self + .model + .lock() + .map_err(|e| MLError::LockError(format!("KAN model lock poisoned: {e}")))?; + let output = model.forward(&input)?; + + // Output: [1, 1] -> squeeze to scalar tensor + let squeezed = output + .squeeze(0) + .and_then(|t| t.squeeze(0)) + .map_err(|e| MLError::ModelError(format!("KAN squeeze: {e}")))?; + + // One sync for confidence (acceptable) + let raw_f32 = squeezed + .to_scalar::() + .map_err(|e| MLError::ModelError(format!("KAN confidence calc: {e}")))?; + let prob = 1.0 / (1.0 + (-raw_f32 as f64).exp()); + let confidence = ((prob - 0.5).abs() * 2.0).clamp(0.0, 1.0); + + Ok(RawPrediction { + direction_scalar: 0.0, + confidence, + tensor: Some(squeezed), + }) + } + fn is_ready(&self) -> bool { self.model.lock().is_ok() } diff --git a/crates/ml/src/ensemble/adapters/tggn.rs b/crates/ml/src/ensemble/adapters/tggn.rs index 953092333..2168f63c1 100644 --- a/crates/ml/src/ensemble/adapters/tggn.rs +++ b/crates/ml/src/ensemble/adapters/tggn.rs @@ -9,7 +9,7 @@ use candle_core::{DType, Device, Module, Tensor}; use candle_nn::{linear, Linear, VarBuilder, VarMap}; use crate::ensemble::inference_adapter::{ - EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, + EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; use crate::gpu::DeviceConfig; use crate::{MLError, MLResult}; @@ -168,6 +168,37 @@ impl ModelInferenceAdapter for TggnInferenceAdapter { }) } + fn predict_raw(&self, features: &FeatureVector) -> MLResult { + let padded = self.pad_features(&features.values); + let input = Tensor::from_vec(padded, (1, self.input_dim), &self.device) + .map_err(|e| MLError::ModelError(format!("TGGN input tensor: {e}")))?; + + let model = self + .model + .lock() + .map_err(|e| MLError::LockError(format!("TGGN lock poisoned: {e}")))?; + let output = model.forward(&input)?; + + // Output: [1, 1] -> squeeze to scalar tensor + let squeezed = output + .squeeze(0) + .and_then(|t| t.squeeze(0)) + .map_err(|e| MLError::ModelError(format!("TGGN squeeze: {e}")))?; + + // One sync for confidence (acceptable) + let raw_f32 = squeezed + .to_scalar::() + .map_err(|e| MLError::ModelError(format!("TGGN confidence calc: {e}")))?; + let prob = 1.0 / (1.0 + (-raw_f32 as f64).exp()); + let confidence = ((prob - 0.5).abs() * 2.0).clamp(0.0, 1.0); + + Ok(RawPrediction { + direction_scalar: 0.0, + confidence, + tensor: Some(squeezed), + }) + } + fn is_ready(&self) -> bool { self.model.lock().is_ok() } diff --git a/crates/ml/src/ensemble/adapters/tlob.rs b/crates/ml/src/ensemble/adapters/tlob.rs index b91d3ac03..066051942 100644 --- a/crates/ml/src/ensemble/adapters/tlob.rs +++ b/crates/ml/src/ensemble/adapters/tlob.rs @@ -13,7 +13,7 @@ use candle_core::{DType, Device, Module, Tensor}; use candle_nn::{linear, Linear, VarBuilder, VarMap}; use crate::ensemble::inference_adapter::{ - EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, + EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; use crate::gpu::DeviceConfig; use crate::{MLError, MLResult}; @@ -240,6 +240,73 @@ impl ModelInferenceAdapter for TlobInferenceAdapter { }) } + fn predict_raw(&self, features: &FeatureVector) -> MLResult { + let padded = self.pad_features(&features.values); + + // Buffer phase (short-lived lock) + let buffer_ready = { + let mut buf = self + .buffer + .lock() + .map_err(|e| MLError::LockError(format!("TLOB buffer lock poisoned: {e}")))?; + buf.push_back(padded); + while buf.len() > self.sequence_length { + buf.pop_front(); + } + buf.len() >= self.sequence_length + }; + + if !buffer_ready { + return Ok(RawPrediction { + direction_scalar: 0.0, + confidence: 0.0, + tensor: None, + }); + } + + // Flatten buffer to [1, seq_len * feature_dim] + let flat_data = { + let buf = self + .buffer + .lock() + .map_err(|e| MLError::LockError(format!("TLOB buffer lock poisoned: {e}")))?; + let mut data = Vec::with_capacity(self.sequence_length * self.feature_dim); + for frame in buf.iter() { + data.extend_from_slice(frame); + } + data + }; + + let flat_dim = self.sequence_length * self.feature_dim; + let input = Tensor::from_vec(flat_data, (1, flat_dim), &self.device) + .map_err(|e| MLError::ModelError(format!("TLOB input tensor: {e}")))?; + + let model = self + .model + .lock() + .map_err(|e| MLError::LockError(format!("TLOB model lock poisoned: {e}")))?; + let output = model.forward(&input)?; + + // Output: [1, 1] -> squeeze to scalar tensor + let squeezed = output + .squeeze(0) + .and_then(|t| t.squeeze(0)) + .map_err(|e| MLError::ModelError(format!("TLOB squeeze: {e}")))?; + + // One sync for confidence (acceptable) + let raw_f32 = squeezed + .to_scalar::() + .map_err(|e| MLError::ModelError(format!("TLOB confidence calc: {e}")))?; + let prob = 1.0 / (1.0 + (-raw_f32 as f64).exp()); + let confidence = ((prob - 0.5).abs() * 2.0).clamp(0.0, 1.0); + + Ok(RawPrediction { + direction_scalar: 0.0, + confidence, + tensor: Some(squeezed), + }) + } + fn is_ready(&self) -> bool { let buf_ok = self .buffer diff --git a/crates/ml/src/ensemble/adapters/xlstm.rs b/crates/ml/src/ensemble/adapters/xlstm.rs index fba860195..6b9a6202d 100644 --- a/crates/ml/src/ensemble/adapters/xlstm.rs +++ b/crates/ml/src/ensemble/adapters/xlstm.rs @@ -12,7 +12,7 @@ use candle_core::{DType, Device, Tensor}; use candle_nn::{VarBuilder, VarMap}; use crate::ensemble::inference_adapter::{ - EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, + EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta, RawPrediction, }; use crate::gpu::DeviceConfig; use crate::xlstm::config::XLSTMConfig; @@ -203,6 +203,76 @@ impl ModelInferenceAdapter for XlstmInferenceAdapter { }) } + fn predict_raw(&self, features: &FeatureVector) -> MLResult { + let padded = self.pad_features(&features.values); + + // Buffer phase (short-lived lock) + let buffer_ready = { + let mut buf = self + .buffer + .lock() + .map_err(|e| MLError::LockError(format!("xLSTM buffer lock poisoned: {e}")))?; + buf.push_back(padded); + while buf.len() > self.sequence_length { + buf.pop_front(); + } + buf.len() >= self.sequence_length + }; + + if !buffer_ready { + return Ok(RawPrediction { + direction_scalar: 0.0, + confidence: 0.0, + tensor: None, + }); + } + + // Build [1, seq_len, input_dim] tensor from buffer snapshot + let flat_data = { + let buf = self + .buffer + .lock() + .map_err(|e| MLError::LockError(format!("xLSTM buffer lock poisoned: {e}")))?; + let mut data = Vec::with_capacity(self.sequence_length * self.input_dim); + for frame in buf.iter() { + data.extend_from_slice(frame); + } + data + }; + + let input = Tensor::from_vec( + flat_data, + (1, self.sequence_length, self.input_dim), + &self.device, + ) + .map_err(|e| MLError::ModelError(format!("xLSTM input tensor: {e}")))?; + + let model = self + .model + .lock() + .map_err(|e| MLError::LockError(format!("xLSTM model lock poisoned: {e}")))?; + let output = model.forward(&input)?; + + // Output: [1, output_dim] -> squeeze to scalar tensor + let squeezed = output + .squeeze(0) + .and_then(|t| t.squeeze(0)) + .map_err(|e| MLError::ModelError(format!("xLSTM squeeze: {e}")))?; + + // One sync for confidence (acceptable) + let raw_f32 = squeezed + .to_scalar::() + .map_err(|e| MLError::ModelError(format!("xLSTM confidence calc: {e}")))?; + let prob = 1.0 / (1.0 + (-raw_f32 as f64).exp()); + let confidence = ((prob - 0.5).abs() * 2.0).clamp(0.0, 1.0); + + Ok(RawPrediction { + direction_scalar: 0.0, + confidence, + tensor: Some(squeezed), + }) + } + fn is_ready(&self) -> bool { let buf_ok = self .buffer