use ruthril::core::statistics::Statistics;
pub struct OnnxRuntime {
model_path: String,
session: Option<String>, }
impl OnnxRuntime {
pub fn new(model_path: &str) -> Self {
Self {
model_path: model_path.to_string(),
session: None,
}
}
pub fn load_model(&mut self) -> Result<(), Box<dyn std::error::Error>> {
println!("Loading ONNX model from: {}", self.model_path);
self.session = Some("mock_session".to_string());
Ok(())
}
pub fn run_inference(&self, input_data: &[f64]) -> Result<Vec<f64>, Box<dyn std::error::Error>> {
if self.session.is_none() {
return Err("Model not loaded".into());
}
if let Some(mean) = Statistics::mean(input_data) {
println!("Input data mean: {}", mean);
}
let output: Vec<f64> = input_data.iter().map(|x| x * 2.0).collect();
if let Some(std) = Statistics::std_deviation(&output) {
println!("Output standard deviation: {}", std);
}
Ok(output)
}
pub fn benchmark(&self, test_data: &[Vec<f64>]) -> Result<BenchmarkResults, Box<dyn std::error::Error>> {
let mut inference_times = Vec::new();
let mut all_outputs = Vec::new();
for input in test_data {
let start = std::time::Instant::now();
let output = self.run_inference(input)?;
let duration = start.elapsed().as_millis() as f64;
inference_times.push(duration);
all_outputs.extend(output);
}
Ok(BenchmarkResults {
mean_inference_time: Statistics::mean(&inference_times).unwrap_or(0.0),
inference_time_std: Statistics::std_deviation(&inference_times).unwrap_or(0.0),
output_mean: Statistics::mean(&all_outputs).unwrap_or(0.0),
output_std: Statistics::std_deviation(&all_outputs).unwrap_or(0.0),
})
}
}
#[derive(Debug)]
pub struct BenchmarkResults {
pub mean_inference_time: f64,
pub inference_time_std: f64,
pub output_mean: f64,
pub output_std: f64,
}