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 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-02 17:46:09 +01:00
parent f8cec1d6f9
commit 78f5ead601
5 changed files with 238 additions and 5 deletions

View File

@@ -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<RawPrediction> {
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()
}

View File

@@ -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<RawPrediction> {
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::<f32>()
.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()
}

View File

@@ -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<RawPrediction> {
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::<f32>()
.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()
}

View File

@@ -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<RawPrediction> {
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::<f32>()
.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

View File

@@ -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<RawPrediction> {
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::<f32>()
.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