use std::path::Path;
use ndarray::{Array, ArrayD, ArrayView, IxDyn};
use oxigdal_core::buffer::RasterBuffer;
use oxigdal_core::types::RasterDataType;
use oxionnx::{GraphOptimizationLevel, Session, SessionBuilder, Tensor};
use serde::{Deserialize, Serialize};
use tracing::{debug, info};
use crate::error::{InferenceError, ModelError, Result};
use crate::models::Model;
pub struct OnnxModel {
session: Session,
metadata: ModelMetadata,
config: SessionConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelMetadata {
pub name: String,
pub version: String,
pub description: String,
pub input_names: Vec<String>,
pub output_names: Vec<String>,
pub input_shape: (usize, usize, usize),
pub output_shape: (usize, usize, usize),
pub class_labels: Option<Vec<String>>,
}
#[derive(Debug, Clone)]
pub struct SessionConfig {
pub execution_provider: ExecutionProvider,
pub num_threads: usize,
pub graph_optimization: bool,
pub batch_size: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ExecutionProvider {
Cpu,
#[cfg(feature = "gpu")]
Cuda,
#[cfg(feature = "coreml")]
CoreMl,
}
impl Default for SessionConfig {
fn default() -> Self {
Self {
execution_provider: ExecutionProvider::Cpu,
num_threads: num_cpus(),
graph_optimization: true,
batch_size: 1,
}
}
}
impl OnnxModel {
pub fn from_file<P: AsRef<Path>>(path: P) -> Result<Self> {
Self::from_file_with_config(path, SessionConfig::default())
}
pub fn from_file_with_config<P: AsRef<Path>>(path: P, config: SessionConfig) -> Result<Self> {
let path = path.as_ref();
info!("Loading ONNX model from: {}", path.display());
if !path.exists() {
return Err(ModelError::NotFound {
path: path.display().to_string(),
}
.into());
}
let mut builder = SessionBuilder::new();
builder = builder.with_intra_threads(config.num_threads);
if config.graph_optimization {
builder = builder.with_optimization_level(GraphOptimizationLevel::All);
}
#[cfg(feature = "gpu")]
{
use oxionnx::CUDAExecutionProvider;
if matches!(config.execution_provider, ExecutionProvider::Cuda) {
builder = builder.with_execution_providers([CUDAExecutionProvider.build()]);
}
}
#[cfg(feature = "coreml")]
{
use oxionnx::CoreMLExecutionProvider;
if matches!(config.execution_provider, ExecutionProvider::CoreMl) {
builder = builder.with_execution_providers([CoreMLExecutionProvider.build()]);
}
}
let session = builder
.commit_from_file(path)
.map_err(|e| ModelError::LoadFailed {
reason: format!("Failed to load ONNX model: {}", e),
})?;
info!("ONNX model loaded successfully");
let metadata = Self::extract_metadata(&session)?;
Ok(Self {
session,
metadata,
config,
})
}
fn extract_metadata(session: &Session) -> Result<ModelMetadata> {
let inputs = session.input_info();
let outputs = session.output_info();
debug!(
"Extracting metadata: {} inputs, {} outputs",
inputs.len(),
outputs.len()
);
let input_names: Vec<String> = inputs.iter().map(|i| i.name.clone()).collect();
let input_shape = if let Some(first_input) = inputs.first() {
let shape = &first_input.shape;
if shape.len() >= 4 {
let c = shape[1].unwrap_or(3);
let h = shape[2].unwrap_or(256);
let w = shape[3].unwrap_or(256);
(c, h, w)
} else if shape.len() == 3 {
let c = shape[0].unwrap_or(3);
let h = shape[1].unwrap_or(256);
let w = shape[2].unwrap_or(256);
(c, h, w)
} else {
(3, 256, 256) }
} else {
return Err(ModelError::LoadFailed {
reason: "No input tensors found in model".to_string(),
}
.into());
};
let output_names: Vec<String> = outputs.iter().map(|o| o.name.clone()).collect();
let output_shape = if let Some(first_output) = outputs.first() {
let shape = &first_output.shape;
if shape.len() >= 4 {
let c = shape[1].unwrap_or(1);
let h = shape[2].unwrap_or(256);
let w = shape[3].unwrap_or(256);
(c, h, w)
} else if shape.len() == 3 {
let c = shape[0].unwrap_or(1);
let h = shape[1].unwrap_or(256);
let w = shape[2].unwrap_or(256);
(c, h, w)
} else {
(1, 256, 256) }
} else {
return Err(ModelError::LoadFailed {
reason: "No output tensors found in model".to_string(),
}
.into());
};
Ok(ModelMetadata {
name: "onnx_model".to_string(),
version: "1.0.0".to_string(),
description: "ONNX Runtime model".to_string(),
input_names,
output_names,
input_shape,
output_shape,
class_labels: None,
})
}
pub fn infer(&mut self, input: &RasterBuffer) -> Result<RasterBuffer> {
debug!(
"Running inference on {}x{} buffer",
input.width(),
input.height()
);
let input_array = self.buffer_to_ndarray(input)?;
let input_name = self
.metadata
.input_names
.first()
.ok_or_else(|| InferenceError::Failed {
reason: "No input tensor name available".to_string(),
})?
.clone();
let input_tensor = Tensor::from_ndarray_view(input_array.view());
let inputs_map = oxionnx::inputs![input_name.as_str() => input_tensor].map_err(|e| {
InferenceError::Failed {
reason: format!("Failed to build inputs map: {}", e),
}
})?;
let outputs = self
.session
.run(&inputs_map)
.map_err(|e| InferenceError::Failed {
reason: format!("ONNX inference failed: {}", e),
})?;
let output_name =
self.metadata
.output_names
.first()
.ok_or_else(|| InferenceError::Failed {
reason: "No output tensor name available".to_string(),
})?;
let output_tensor = outputs.get(output_name.as_str()).ok_or_else(|| {
InferenceError::OutputParsingFailed {
reason: format!("Output tensor '{}' not found", output_name),
}
})?;
let output_array = output_tensor.try_extract_array::<f32>().map_err(|e| {
InferenceError::OutputParsingFailed {
reason: format!("Failed to extract output tensor: {}", e),
}
})?;
let output_owned = output_array.to_owned();
drop(outputs);
let output_view = output_owned.view().into_dyn();
self.ndarray_to_buffer(&output_view)
}
pub fn infer_batch(&mut self, inputs: &[RasterBuffer]) -> Result<Vec<RasterBuffer>> {
if inputs.is_empty() {
return Ok(Vec::new());
}
debug!("Running batch inference on {} inputs", inputs.len());
let mut results = Vec::with_capacity(inputs.len());
for input in inputs {
let output = self.infer(input)?;
results.push(output);
}
Ok(results)
}
fn buffer_to_ndarray(&self, buffer: &RasterBuffer) -> Result<ArrayD<f32>> {
let width = buffer.width() as usize;
let height = buffer.height() as usize;
let (channels, expected_height, expected_width) = self.metadata.input_shape;
if width != expected_width || height != expected_height {
return Err(InferenceError::InvalidInputShape {
expected: vec![channels, expected_height, expected_width],
actual: vec![channels, height, width],
}
.into());
}
let data = match buffer.data_type() {
RasterDataType::Float32 => {
let slice = buffer
.as_slice::<f32>()
.map_err(crate::error::MlError::OxiGdal)?;
slice.to_vec()
}
RasterDataType::UInt8 => {
let slice = buffer
.as_slice::<u8>()
.map_err(crate::error::MlError::OxiGdal)?;
slice.iter().map(|&v| f32::from(v) / 255.0).collect()
}
RasterDataType::Int16 => {
let slice = buffer
.as_slice::<i16>()
.map_err(crate::error::MlError::OxiGdal)?;
slice.iter().map(|&v| v as f32).collect()
}
RasterDataType::UInt16 => {
let slice = buffer
.as_slice::<u16>()
.map_err(crate::error::MlError::OxiGdal)?;
slice.iter().map(|&v| f32::from(v) / 65535.0).collect()
}
RasterDataType::Float64 => {
let slice = buffer
.as_slice::<f64>()
.map_err(crate::error::MlError::OxiGdal)?;
slice.iter().map(|&v| v as f32).collect()
}
_ => {
return Err(InferenceError::Failed {
reason: format!("Unsupported data type: {:?}", buffer.data_type()),
}
.into());
}
};
let total_pixels = height * width;
let num_bands = data.len() / total_pixels;
let shape = IxDyn(&[1, num_bands, height, width]);
Array::from_shape_vec(shape, data).map_err(|e| {
InferenceError::Failed {
reason: format!("Failed to create ndarray from buffer: {}", e),
}
.into()
})
}
fn ndarray_to_buffer(&self, array: &ArrayView<f32, IxDyn>) -> Result<RasterBuffer> {
let shape = array.shape();
debug!("Converting ndarray with shape {:?} to RasterBuffer", shape);
let (height, width) = if shape.len() == 4 {
(shape[2], shape[3])
} else if shape.len() == 3 {
(shape[1], shape[2])
} else if shape.len() == 2 {
(shape[0], shape[1])
} else {
return Err(InferenceError::OutputParsingFailed {
reason: format!("Unexpected output shape: {:?}", shape),
}
.into());
};
let data: Vec<f32> = array.iter().copied().collect();
let bytes: Vec<u8> = data.iter().flat_map(|&f: &f32| f.to_le_bytes()).collect();
RasterBuffer::new(
bytes,
width as u64,
height as u64,
RasterDataType::Float32,
oxigdal_core::types::NoDataValue::None,
)
.map_err(crate::error::MlError::OxiGdal)
}
}
impl Model for OnnxModel {
fn metadata(&self) -> &ModelMetadata {
&self.metadata
}
fn predict(&mut self, input: &RasterBuffer) -> Result<RasterBuffer> {
self.infer(input)
}
fn predict_batch(&mut self, inputs: &[RasterBuffer]) -> Result<Vec<RasterBuffer>> {
self.infer_batch(inputs)
}
fn input_shape(&self) -> (usize, usize, usize) {
self.metadata.input_shape
}
fn output_shape(&self) -> (usize, usize, usize) {
self.metadata.output_shape
}
}
fn num_cpus() -> usize {
std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_session_config_default() {
let config = SessionConfig::default();
assert_eq!(config.execution_provider, ExecutionProvider::Cpu);
assert!(config.graph_optimization);
assert_eq!(config.batch_size, 1);
}
#[test]
fn test_metadata_serialization() {
let metadata = ModelMetadata {
name: "test_model".to_string(),
version: "1.0.0".to_string(),
description: "Test model".to_string(),
input_names: vec!["input".to_string()],
output_names: vec!["output".to_string()],
input_shape: (3, 256, 256),
output_shape: (1, 256, 256),
class_labels: None,
};
let json = serde_json::to_string(&metadata);
assert!(json.is_ok());
}
#[test]
fn test_num_cpus() {
let cpus = num_cpus();
assert!(cpus > 0);
assert!(cpus <= 256); }
fn build_session_from_graph(
graph: oxionnx::Graph,
weights: std::collections::HashMap<String, Tensor>,
) -> Result<Session> {
let session = SessionBuilder::new()
.build_from_graph(graph, weights)
.map_err(|e| ModelError::LoadFailed {
reason: format!("Failed to build session from graph: {}", e),
})?;
Ok(session)
}
fn build_identity_graph(
input_name: &str,
output_name: &str,
shape: &[Option<usize>],
) -> oxionnx::Graph {
use oxionnx::{Attributes, DType, Node, OpKind, TensorInfo};
oxionnx::Graph {
name: "identity_test".to_string(),
nodes: vec![Node {
op: OpKind::Identity,
name: "identity_0".to_string(),
inputs: vec![input_name.to_string()],
outputs: vec![output_name.to_string()],
attrs: Attributes::default(),
}],
input_names: vec![input_name.to_string()],
output_names: vec![output_name.to_string()],
input_infos: vec![TensorInfo {
name: input_name.to_string(),
dtype: DType::F32,
shape: shape.to_vec(),
dim_params: vec![],
}],
output_infos: vec![TensorInfo {
name: output_name.to_string(),
dtype: DType::F32,
shape: shape.to_vec(),
dim_params: vec![],
}],
}
}
fn build_relu_graph(
input_name: &str,
output_name: &str,
shape: &[Option<usize>],
) -> oxionnx::Graph {
use oxionnx::{Attributes, DType, Node, OpKind, TensorInfo};
oxionnx::Graph {
name: "relu_test".to_string(),
nodes: vec![Node {
op: OpKind::Relu,
name: "relu_0".to_string(),
inputs: vec![input_name.to_string()],
outputs: vec![output_name.to_string()],
attrs: Attributes::default(),
}],
input_names: vec![input_name.to_string()],
output_names: vec![output_name.to_string()],
input_infos: vec![TensorInfo {
name: input_name.to_string(),
dtype: DType::F32,
shape: shape.to_vec(),
dim_params: vec![],
}],
output_infos: vec![TensorInfo {
name: output_name.to_string(),
dtype: DType::F32,
shape: shape.to_vec(),
dim_params: vec![],
}],
}
}
fn build_add_bias_graph(
input_name: &str,
bias_name: &str,
output_name: &str,
shape: &[Option<usize>],
) -> oxionnx::Graph {
use oxionnx::{Attributes, DType, Node, OpKind, TensorInfo};
oxionnx::Graph {
name: "add_bias_test".to_string(),
nodes: vec![Node {
op: OpKind::Add,
name: "add_0".to_string(),
inputs: vec![input_name.to_string(), bias_name.to_string()],
outputs: vec![output_name.to_string()],
attrs: Attributes::default(),
}],
input_names: vec![input_name.to_string()],
output_names: vec![output_name.to_string()],
input_infos: vec![TensorInfo {
name: input_name.to_string(),
dtype: DType::F32,
shape: shape.to_vec(),
dim_params: vec![],
}],
output_infos: vec![TensorInfo {
name: output_name.to_string(),
dtype: DType::F32,
shape: shape.to_vec(),
dim_params: vec![],
}],
}
}
#[test]
fn test_identity_inference_end_to_end() {
let shape = &[Some(1), Some(3), Some(4), Some(4)];
let graph = build_identity_graph("X", "Y", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build identity session");
let input_data: Vec<f32> = (0..48).map(|i| i as f32 * 0.1).collect();
let input_tensor = Tensor::new(input_data.clone(), vec![1, 3, 4, 4]);
let inputs_map = oxionnx::inputs!["X" => input_tensor].expect("build inputs map");
let outputs = session.run(&inputs_map).expect("run identity inference");
let output = outputs.get("Y").expect("output Y not found");
let (out_shape, out_data) = output
.try_extract_tensor::<f32>()
.expect("extract output tensor");
assert_eq!(out_shape, &[1, 3, 4, 4]);
assert_eq!(out_data.len(), 48);
for (a, b) in input_data.iter().zip(out_data.iter()) {
assert!((a - b).abs() < 1e-6, "identity mismatch: {} vs {}", a, b);
}
}
#[test]
fn test_relu_inference_end_to_end() {
let shape = &[Some(1), Some(1), Some(2), Some(3)];
let graph = build_relu_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build relu session");
let input_data: Vec<f32> = vec![-3.0, -1.0, 0.0, 1.0, 2.5, -0.5];
let expected: Vec<f32> = vec![0.0, 0.0, 0.0, 1.0, 2.5, 0.0];
let input_tensor = Tensor::new(input_data, vec![1, 1, 2, 3]);
let inputs_map = oxionnx::inputs!["input" => input_tensor].expect("build inputs map");
let outputs = session.run(&inputs_map).expect("run relu inference");
let output = outputs.get("output").expect("output not found");
let (out_shape, out_data) = output.try_extract_tensor::<f32>().expect("extract output");
assert_eq!(out_shape, &[1, 1, 2, 3]);
for (a, b) in expected.iter().zip(out_data.iter()) {
assert!(
(a - b).abs() < 1e-6,
"relu mismatch: expected {} got {}",
a,
b
);
}
}
#[test]
fn test_add_bias_inference_end_to_end() {
let shape = &[Some(1), Some(1), Some(2), Some(2)];
let graph = build_add_bias_graph("input", "bias", "output", shape);
let mut weights = std::collections::HashMap::new();
weights.insert(
"bias".to_string(),
Tensor::new(vec![10.0, 20.0, 30.0, 40.0], vec![1, 1, 2, 2]),
);
let session = build_session_from_graph(graph, weights).expect("build add session");
let input_data: Vec<f32> = vec![1.0, 2.0, 3.0, 4.0];
let expected: Vec<f32> = vec![11.0, 22.0, 33.0, 44.0];
let input_tensor = Tensor::new(input_data, vec![1, 1, 2, 2]);
let inputs_map = oxionnx::inputs!["input" => input_tensor].expect("build inputs map");
let outputs = session.run(&inputs_map).expect("run add inference");
let output = outputs.get("output").expect("output not found");
let (out_shape, out_data) = output.try_extract_tensor::<f32>().expect("extract output");
assert_eq!(out_shape, &[1, 1, 2, 2]);
for (a, b) in expected.iter().zip(out_data.iter()) {
assert!(
(a - b).abs() < 1e-6,
"add mismatch: expected {} got {}",
a,
b
);
}
}
#[test]
fn test_metadata_extraction_nchw_shape() {
let shape = &[Some(1), Some(3), Some(64), Some(64)];
let graph = build_identity_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build session for metadata extraction");
let metadata = OnnxModel::extract_metadata(&session).expect("extract metadata");
assert_eq!(metadata.input_names, vec!["input"]);
assert_eq!(metadata.output_names, vec!["output"]);
assert_eq!(metadata.input_shape, (3, 64, 64));
assert_eq!(metadata.output_shape, (3, 64, 64));
}
#[test]
fn test_metadata_extraction_dynamic_dims() {
let shape = &[None, Some(3), Some(32), Some(32)];
let graph = build_identity_graph("x", "y", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build session for dynamic dim test");
let metadata =
OnnxModel::extract_metadata(&session).expect("extract metadata with dynamic dims");
assert_eq!(metadata.input_shape, (3, 32, 32));
assert_eq!(metadata.output_shape, (3, 32, 32));
}
#[test]
fn test_metadata_extraction_3d_shape() {
let shape = &[Some(3), Some(128), Some(128)];
let graph = build_identity_graph("img", "out", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build session for 3D shape test");
let metadata = OnnxModel::extract_metadata(&session).expect("extract 3D metadata");
assert_eq!(metadata.input_shape, (3, 128, 128));
}
#[test]
fn test_session_builder_with_intra_threads() {
let shape = &[Some(1), Some(1), Some(2), Some(2)];
let graph = build_identity_graph("x", "y", shape);
let session = SessionBuilder::new()
.with_intra_threads(2)
.build_from_graph(graph, std::collections::HashMap::new())
.map_err(|e| ModelError::LoadFailed {
reason: e.to_string(),
})
.expect("build session with intra_threads");
let input_tensor = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], vec![1, 1, 2, 2]);
let inputs_map = oxionnx::inputs!["x" => input_tensor].expect("build inputs map");
let outputs = session.run(&inputs_map).expect("run with intra_threads");
assert!(outputs.contains_key("y"));
}
#[test]
fn test_execution_provider_variants() {
let cpu = ExecutionProvider::Cpu;
assert_eq!(cpu, ExecutionProvider::Cpu);
#[cfg(feature = "gpu")]
{
let cuda = ExecutionProvider::Cuda;
assert_eq!(cuda, ExecutionProvider::Cuda);
assert_ne!(cuda, ExecutionProvider::Cpu);
}
}
#[test]
fn test_two_node_pipeline_relu_identity() {
use oxionnx::{Attributes, DType, Node, OpKind, TensorInfo};
let graph = oxionnx::Graph {
name: "relu_identity_pipeline".to_string(),
nodes: vec![
Node {
op: OpKind::Relu,
name: "relu_0".to_string(),
inputs: vec!["input".to_string()],
outputs: vec!["intermediate".to_string()],
attrs: Attributes::default(),
},
Node {
op: OpKind::Identity,
name: "identity_0".to_string(),
inputs: vec!["intermediate".to_string()],
outputs: vec!["output".to_string()],
attrs: Attributes::default(),
},
],
input_names: vec!["input".to_string()],
output_names: vec!["output".to_string()],
input_infos: vec![TensorInfo {
name: "input".to_string(),
dtype: DType::F32,
shape: vec![Some(1), Some(1), Some(2), Some(3)],
dim_params: vec![],
}],
output_infos: vec![TensorInfo {
name: "output".to_string(),
dtype: DType::F32,
shape: vec![Some(1), Some(1), Some(2), Some(3)],
dim_params: vec![],
}],
};
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build pipeline session");
let input_data: Vec<f32> = vec![-5.0, -1.0, 0.0, 3.0, 7.0, -2.0];
let expected: Vec<f32> = vec![0.0, 0.0, 0.0, 3.0, 7.0, 0.0];
let input_tensor = Tensor::new(input_data, vec![1, 1, 2, 3]);
let inputs_map = oxionnx::inputs!["input" => input_tensor].expect("build inputs map");
let outputs = session.run(&inputs_map).expect("run pipeline");
let output = outputs.get("output").expect("output not found");
let (_shape, out_data) = output.try_extract_tensor::<f32>().expect("extract output");
for (a, b) in expected.iter().zip(out_data.iter()) {
assert!(
(a - b).abs() < 1e-6,
"pipeline mismatch: expected {} got {}",
a,
b,
);
}
}
#[test]
fn test_ndarray_tensor_roundtrip() {
let arr = ndarray::Array::from_shape_vec(
ndarray::IxDyn(&[1, 2, 3]),
vec![1.0_f32, 2.0, 3.0, 4.0, 5.0, 6.0],
)
.expect("create ndarray");
let tensor = Tensor::from_ndarray_view(arr.view());
let extracted = tensor
.try_extract_array::<f32>()
.expect("extract array from tensor");
assert_eq!(extracted.shape(), &[1, 2, 3]);
for (a, b) in arr.iter().zip(extracted.iter()) {
assert!((a - b).abs() < 1e-6);
}
}
#[test]
fn test_model_not_found_error() {
let result = OnnxModel::from_file("/nonexistent/path/model.onnx");
assert!(result.is_err());
let err_msg = format!("{}", result.err().expect("should be error"));
assert!(
err_msg.contains("not found") || err_msg.contains("Not"),
"error should mention 'not found', got: {}",
err_msg,
);
}
#[test]
fn test_gpu_feature_compilation() {
let config = SessionConfig {
execution_provider: ExecutionProvider::Cpu,
num_threads: 1,
graph_optimization: false,
batch_size: 2,
};
assert_eq!(config.batch_size, 2);
assert!(!config.graph_optimization);
}
}