use std::path::Path;
use ndarray::{Array, ArrayD, ArrayView, IxDyn};
use oxigdal_core::buffer::{MultiBandBuffer, RasterBuffer};
use oxigdal_core::types::{ColorInterpretation, 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 output_owned = self.run_forward(input_array)?;
let output_view = output_owned.view();
self.ndarray_to_buffer(&output_view)
}
pub fn infer_multiband(&mut self, input: &MultiBandBuffer) -> Result<MultiBandBuffer> {
debug!(
"Running multi-band inference on {}x{} buffer with {} band(s)",
input.width(),
input.height(),
input.band_count()
);
let input_array = self.multiband_buffer_to_ndarray(input)?;
let output_owned = self.run_forward(input_array)?;
let output_view = output_owned.view();
self.ndarray_to_multiband_buffer(&output_view)
}
fn run_forward(&mut self, input_array: ArrayD<f32>) -> Result<ArrayD<f32>> {
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().into_dyn();
drop(outputs);
Ok(output_owned)
}
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)
}
pub fn infer_batch_multiband(
&mut self,
inputs: &[MultiBandBuffer],
) -> Result<Vec<MultiBandBuffer>> {
if inputs.is_empty() {
return Ok(Vec::new());
}
debug!(
"Running multi-band batch inference on {} inputs",
inputs.len()
);
let mut results = Vec::with_capacity(inputs.len());
for input in inputs {
results.push(self.infer_multiband(input)?);
}
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 = band_buffer_to_f32(buffer)?;
let total_pixels = height * width;
if total_pixels == 0 {
return Err(InferenceError::Failed {
reason: "Input buffer has zero pixels".to_string(),
}
.into());
}
let num_bands = data.len() / total_pixels;
if num_bands != channels {
return Err(InferenceError::InvalidBandCount {
expected: channels,
actual: num_bands,
}
.into());
}
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)
}
fn multiband_buffer_to_ndarray(&self, buffer: &MultiBandBuffer) -> Result<ArrayD<f32>> {
let width = buffer.width() as usize;
let height = buffer.height() as usize;
let num_bands = buffer.band_count() 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![num_bands, height, width],
}
.into());
}
if num_bands != channels {
return Err(InferenceError::InvalidBandCount {
expected: channels,
actual: num_bands,
}
.into());
}
let total_pixels = height * width;
if total_pixels == 0 || num_bands == 0 {
return Err(InferenceError::Failed {
reason: "Input buffer has zero pixels or zero bands".to_string(),
}
.into());
}
let mut data = Vec::with_capacity(num_bands * total_pixels);
for b in 0..num_bands as u32 {
let band = buffer.band(b).map_err(crate::error::MlError::OxiGdal)?;
let band_data = band_buffer_to_f32(band.buffer())?;
if band_data.len() != total_pixels {
return Err(InferenceError::Failed {
reason: format!(
"Band {} decoded to {} pixels, expected {}",
b,
band_data.len(),
total_pixels
),
}
.into());
}
data.extend_from_slice(&band_data);
}
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 multi-band buffer: {}", e),
}
.into()
})
}
fn ndarray_to_multiband_buffer(
&self,
array: &ArrayView<f32, IxDyn>,
) -> Result<MultiBandBuffer> {
let shape = array.shape();
debug!(
"Converting ndarray with shape {:?} to MultiBandBuffer",
shape
);
let (channels, height, width) = match shape.len() {
4 => (shape[1], shape[2], shape[3]),
3 => (shape[0], shape[1], shape[2]),
2 => (1, shape[0], shape[1]),
_ => {
return Err(InferenceError::OutputParsingFailed {
reason: format!("Unexpected output shape: {:?}", shape),
}
.into());
}
};
if channels == 0 || height == 0 || width == 0 {
return Err(InferenceError::OutputParsingFailed {
reason: format!("Output shape has a zero dimension: {:?}", shape),
}
.into());
}
let data: Vec<f32> = array.iter().copied().collect();
let per_band = height * width;
let expected = channels * per_band;
if data.len() != expected {
return Err(InferenceError::OutputParsingFailed {
reason: format!(
"Output element count {} does not match C*H*W = {}",
data.len(),
expected
),
}
.into());
}
let mut bands = Vec::with_capacity(channels);
for c in 0..channels {
let start = c * per_band;
let band_data = data[start..start + per_band].to_vec();
let band =
RasterBuffer::from_typed_vec(width, height, band_data, RasterDataType::Float32)
.map_err(crate::error::MlError::OxiGdal)?;
bands.push(band);
}
let colors = default_band_colors(channels);
MultiBandBuffer::from_bands(bands, colors).map_err(crate::error::MlError::OxiGdal)
}
}
fn band_buffer_to_f32(buffer: &RasterBuffer) -> Result<Vec<f32>> {
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()
}
other => {
return Err(InferenceError::Failed {
reason: format!("Unsupported data type: {:?}", other),
}
.into());
}
};
Ok(data)
}
fn default_band_colors(n: usize) -> Vec<ColorInterpretation> {
match n {
1 => vec![ColorInterpretation::Gray],
3 => vec![
ColorInterpretation::Red,
ColorInterpretation::Green,
ColorInterpretation::Blue,
],
4 => vec![
ColorInterpretation::Red,
ColorInterpretation::Green,
ColorInterpretation::Blue,
ColorInterpretation::Alpha,
],
_ => vec![ColorInterpretation::Undefined; n],
}
}
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 predict_multiband(&mut self, input: &MultiBandBuffer) -> Result<MultiBandBuffer> {
self.infer_multiband(input)
}
fn predict_batch_multiband(
&mut self,
inputs: &[MultiBandBuffer],
) -> Result<Vec<MultiBandBuffer>> {
self.infer_batch_multiband(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_buffer_to_ndarray_rejects_band_count_mismatch() {
use crate::error::MlError;
use oxigdal_core::buffer::RasterBuffer;
use oxigdal_core::types::RasterDataType;
let shape = &[Some(1), Some(3), Some(4), Some(4)];
let graph = build_identity_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build 3-channel session");
let metadata = OnnxModel::extract_metadata(&session).expect("extract metadata");
assert_eq!(metadata.input_shape, (3, 4, 4));
let mut model = OnnxModel {
session,
metadata,
config: SessionConfig::default(),
};
let buffer = RasterBuffer::zeros(4, 4, RasterDataType::Float32);
let result = model.infer(&buffer);
assert!(matches!(
result,
Err(MlError::Inference(InferenceError::InvalidBandCount {
expected: 3,
actual: 1,
}))
));
}
#[test]
fn test_buffer_to_ndarray_accepts_single_band_model() {
use oxigdal_core::buffer::RasterBuffer;
use oxigdal_core::types::RasterDataType;
let shape = &[Some(1), Some(1), Some(4), Some(4)];
let graph = build_identity_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build 1-channel session");
let metadata = OnnxModel::extract_metadata(&session).expect("extract metadata");
assert_eq!(metadata.input_shape, (1, 4, 4));
let mut model = OnnxModel {
session,
metadata,
config: SessionConfig::default(),
};
let buffer = RasterBuffer::zeros(4, 4, RasterDataType::Float32);
let result = model.infer(&buffer);
assert!(
result.is_ok(),
"single-band inference failed: {:?}",
result.err()
);
}
fn make_multiband(
width: usize,
height: usize,
values: &[f32],
) -> oxigdal_core::buffer::MultiBandBuffer {
use oxigdal_core::buffer::{MultiBandBuffer, RasterBuffer};
use oxigdal_core::types::{ColorInterpretation, RasterDataType};
let per_band = width * height;
let mut bands = Vec::with_capacity(values.len());
for &v in values {
let data = vec![v; per_band];
let band = RasterBuffer::from_typed_vec(width, height, data, RasterDataType::Float32)
.expect("build float32 band");
bands.push(band);
}
let colors = vec![ColorInterpretation::Undefined; values.len()];
MultiBandBuffer::from_bands(bands, colors).expect("build multi-band buffer")
}
#[test]
fn test_multiband_buffer_to_ndarray_layout() {
let shape = &[Some(1), Some(3), Some(2), Some(4)];
let graph = build_identity_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build 3-channel session");
let metadata = OnnxModel::extract_metadata(&session).expect("extract metadata");
assert_eq!(metadata.input_shape, (3, 2, 4));
let model = OnnxModel {
session,
metadata,
config: SessionConfig::default(),
};
let buffer = make_multiband(4, 2, &[10.0, 20.0, 30.0]);
let tensor = model
.multiband_buffer_to_ndarray(&buffer)
.expect("convert multi-band buffer to tensor");
assert_eq!(tensor.shape(), &[1, 3, 2, 4]);
let per_band = 2 * 4;
let flat: Vec<f32> = tensor.iter().copied().collect();
assert_eq!(flat.len(), 3 * per_band);
for i in 0..per_band {
assert!((flat[i] - 10.0).abs() < 1e-6, "band0 elem {}", i);
assert!((flat[per_band + i] - 20.0).abs() < 1e-6, "band1 elem {}", i);
assert!(
(flat[2 * per_band + i] - 30.0).abs() < 1e-6,
"band2 elem {}",
i
);
}
}
#[test]
fn test_multiband_buffer_to_ndarray_rejects_band_count_mismatch() {
use crate::error::MlError;
let shape = &[Some(1), Some(3), Some(2), Some(2)];
let graph = build_identity_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build 3-channel session");
let metadata = OnnxModel::extract_metadata(&session).expect("extract metadata");
let model = OnnxModel {
session,
metadata,
config: SessionConfig::default(),
};
let buffer = make_multiband(2, 2, &[1.0, 2.0]); let result = model.multiband_buffer_to_ndarray(&buffer);
assert!(matches!(
result,
Err(MlError::Inference(InferenceError::InvalidBandCount {
expected: 3,
actual: 2,
}))
));
}
#[test]
fn test_ndarray_to_multiband_buffer_roundtrip() {
let shape = &[Some(1), Some(1), Some(2), Some(2)];
let graph = build_identity_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build session");
let metadata = OnnxModel::extract_metadata(&session).expect("extract metadata");
let model = OnnxModel {
session,
metadata,
config: SessionConfig::default(),
};
let data: Vec<f32> = (1..=12).map(|i| i as f32).collect();
let arr = ndarray::Array::from_shape_vec(ndarray::IxDyn(&[1, 3, 2, 2]), data)
.expect("build tensor");
let view = arr.view();
let multi = model
.ndarray_to_multiband_buffer(&view)
.expect("unpack tensor to multi-band buffer");
assert_eq!(multi.band_count(), 3);
assert_eq!(multi.width(), 2);
assert_eq!(multi.height(), 2);
let expected: [[f64; 4]; 3] = [
[1.0, 2.0, 3.0, 4.0],
[5.0, 6.0, 7.0, 8.0],
[9.0, 10.0, 11.0, 12.0],
];
for b in 0..3u32 {
let band = multi.band(b).expect("band ref");
let buf = band.buffer();
for y in 0..2u64 {
for x in 0..2u64 {
let got = buf.get_pixel(x, y).expect("pixel");
let want = expected[b as usize][(y * 2 + x) as usize];
assert!((got - want).abs() < 1e-6, "band {} ({},{})", b, x, y);
}
}
}
}
#[test]
fn test_multiband_identity_inference_end_to_end() {
let shape = &[Some(1), Some(3), Some(2), Some(2)];
let graph = build_identity_graph("input", "output", shape);
let session = build_session_from_graph(graph, std::collections::HashMap::new())
.expect("build 3-channel identity session");
let metadata = OnnxModel::extract_metadata(&session).expect("extract metadata");
let mut model = OnnxModel {
session,
metadata,
config: SessionConfig::default(),
};
let buffer = make_multiband(2, 2, &[7.0, 8.0, 9.0]);
let out = model
.infer_multiband(&buffer)
.expect("multi-band inference");
assert_eq!(out.band_count(), 3);
assert_eq!(out.width(), 2);
assert_eq!(out.height(), 2);
let expected = [7.0f64, 8.0, 9.0];
for b in 0..3u32 {
let band = out.band(b).expect("band ref");
let buf = band.buffer();
for y in 0..2u64 {
for x in 0..2u64 {
let got = buf.get_pixel(x, y).expect("pixel");
assert!(
(got - expected[b as usize]).abs() < 1e-6,
"band {} ({},{}) got {}",
b,
x,
y,
got
);
}
}
}
}
#[test]
fn test_default_band_colors_mapping() {
use oxigdal_core::types::ColorInterpretation;
assert_eq!(default_band_colors(1), vec![ColorInterpretation::Gray]);
assert_eq!(
default_band_colors(3),
vec![
ColorInterpretation::Red,
ColorInterpretation::Green,
ColorInterpretation::Blue
]
);
assert_eq!(default_band_colors(2).len(), 2);
assert!(
default_band_colors(5)
.iter()
.all(|c| *c == ColorInterpretation::Undefined)
);
}
#[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);
}
}