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:
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user