use std::sync::Mutex;
use ndarray::{ArrayD, IxDyn};
use ort::session::Session;
use ort::session::builder::GraphOptimizationLevel;
use ort::value::Tensor as OrtTensor;
use super::{ModelBackend, Tensor};
use crate::error::{OcrError, Result};
pub(crate) struct OrtBackend {
session: Mutex<Session>,
}
impl OrtBackend {
pub(crate) fn load(model_bytes: &[u8], threads: usize) -> Result<Self> {
let mut builder =
Session::builder().map_err(|error| inference_error("create an ONNX Runtime session builder", error))?;
builder = builder
.with_optimization_level(GraphOptimizationLevel::All)
.map_err(|error| {
inference_error("set the ONNX Runtime graph optimization level", ort::Error::from(error))
})?;
builder = builder.with_memory_pattern(false).map_err(|error| {
inference_error("disable ONNX Runtime memory-pattern planning", ort::Error::from(error))
})?;
if threads > 0 {
builder = builder
.with_intra_threads(threads)
.map_err(|error| inference_error("configure ONNX Runtime intra-op threads", ort::Error::from(error)))?;
}
let session = builder
.commit_from_memory(model_bytes)
.map_err(|error| inference_error("load the ONNX model from memory", error))?;
Ok(Self {
session: Mutex::new(session),
})
}
}
impl ModelBackend for OrtBackend {
fn name(&self) -> &str {
"ort"
}
fn run(&self, input: Tensor) -> Result<Tensor> {
let (shape, data) = input_buffer(input);
let value = OrtTensor::from_array((shape, data))
.map_err(|error| inference_error("build the ONNX Runtime input tensor", error))?;
let mut session = self.session.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
let outputs = session
.run(ort::inputs![value])
.map_err(|error| inference_error("run ONNX Runtime inference", error))?;
let output = outputs
.values()
.next()
.ok_or_else(|| OcrError::inference("ONNX Runtime returned no output tensor"))?;
let (out_shape, out_data) = output
.try_extract_tensor::<f32>()
.map_err(|error| inference_error("extract the ONNX Runtime output tensor", error))?;
array_from_output(out_shape, out_data)
}
}
fn input_buffer(input: Tensor) -> (Vec<i64>, Vec<f32>) {
let shape = shape_to_i64(input.shape());
let element_count: usize = input.shape().iter().product();
if input.is_standard_layout() {
let (data, offset) = input.into_raw_vec_and_offset();
let start = offset.unwrap_or(0);
if start == 0 && data.len() == element_count {
return (shape, data);
}
return (shape, data.into_iter().skip(start).take(element_count).collect());
}
let (data, _) = input.as_standard_layout().into_owned().into_raw_vec_and_offset();
(shape, data)
}
fn shape_to_i64(shape: &[usize]) -> Vec<i64> {
shape.iter().map(|&dim| dim as i64).collect()
}
fn array_from_output(dims: &[i64], data: &[f32]) -> Result<Tensor> {
if let Some(&negative) = dims.iter().find(|&&dim| dim < 0) {
return Err(OcrError::inference(format!(
"ONNX Runtime returned a negative output dimension {negative} in shape {dims:?}"
)));
}
let shape: Vec<usize> = dims.iter().map(|&dim| dim as usize).collect();
ArrayD::from_shape_vec(IxDyn(&shape), data.to_vec()).map_err(|error| OcrError::Inference {
message: format!(
"ONNX Runtime output shape {dims:?} does not match {} elements",
data.len()
),
source: Some(Box::new(error)),
})
}
fn inference_error(operation: &str, source: ort::Error) -> OcrError {
OcrError::Inference {
message: format!("ONNX Runtime backend failed to {operation}"),
source: Some(Box::new(source)),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn shape_to_i64_converts_dims() {
assert_eq!(shape_to_i64(&[1, 3, 64, 128]), vec![1_i64, 3, 64, 128]);
}
#[test]
fn shape_to_i64_handles_empty_shape() {
assert_eq!(shape_to_i64(&[]), Vec::<i64>::new());
}
#[test]
fn array_from_output_preserves_shape() {
let data = [1.0_f32, 2.0, 3.0, 4.0, 5.0, 6.0];
let array = array_from_output(&[1, 2, 3], &data).expect("shape matches element count");
assert_eq!(array.shape(), &[1, 2, 3]);
assert_eq!(array.iter().copied().collect::<Vec<_>>(), data);
}
#[test]
fn array_from_output_errors_on_length_mismatch() {
let error = array_from_output(&[2, 2], &[1.0_f32, 2.0, 3.0]).expect_err("length mismatch must fail");
assert!(matches!(error, OcrError::Inference { .. }));
}
#[test]
fn input_buffer_preserves_row_major_order_for_standard_layout() {
let expected: Vec<f32> = (0..8).map(|value| value as f32).collect();
let input = ArrayD::from_shape_vec(IxDyn(&[1, 2, 2, 2]), expected.clone()).expect("build the tensor");
let manual: Vec<f32> = input.iter().copied().collect();
let (shape, data) = input_buffer(input);
assert_eq!(shape, vec![1_i64, 2, 2, 2]);
assert_eq!(data, manual);
assert_eq!(data, expected);
}
#[test]
fn input_buffer_matches_iter_order_for_non_standard_layout() {
let base = ArrayD::from_shape_vec(IxDyn(&[2, 3]), (0..6).map(|value| value as f32).collect())
.expect("build the tensor");
let transposed = base.t().into_owned();
assert!(
!transposed.is_standard_layout(),
"transpose must be non-standard layout"
);
let manual: Vec<f32> = transposed.iter().copied().collect();
let (shape, data) = input_buffer(transposed);
assert_eq!(shape, vec![3_i64, 2]);
assert_eq!(data, manual);
}
#[test]
fn input_buffer_slices_standard_layout_with_nonzero_offset() {
use ndarray::{Axis, Slice};
let mut base = ArrayD::from_shape_vec(IxDyn(&[3, 2]), (0..6).map(|value| value as f32).collect())
.expect("build the tensor");
base.slice_axis_inplace(Axis(0), Slice::from(1..));
assert!(base.is_standard_layout(), "the sliced array must stay standard layout");
let manual: Vec<f32> = base.iter().copied().collect();
let (shape, data) = input_buffer(base);
assert_eq!(shape, vec![2_i64, 2]);
assert_eq!(
data, manual,
"must return the logical elements, not the full backing buffer"
);
}
#[test]
fn array_from_output_errors_on_negative_dimension() {
let error = array_from_output(&[-1, 2], &[1.0_f32, 2.0]).expect_err("negative dimension must fail");
assert!(matches!(error, OcrError::Inference { .. }));
}
#[test]
#[ignore = "requires the ONNX Runtime native library and a model file"]
fn load_and_run_over_real_model() {
let model_path = std::env::var("EASYOCR_TEST_ONNX").expect("set EASYOCR_TEST_ONNX to an ONNX model path");
let model_bytes = std::fs::read(&model_path).expect("read the model file");
let backend = OrtBackend::load(&model_bytes, 1).expect("load the ONNX model");
assert_eq!(backend.name(), "ort");
let input = ArrayD::from_elem(IxDyn(&[1, 3, 64, 64]), 1.0_f32);
let output = backend.run(input).expect("run inference");
assert!(output.ndim() >= 2, "expected a multi-dimensional output");
assert!(!output.is_empty(), "expected a non-empty output");
assert!(
output.iter().all(|value| value.is_finite()),
"expected all outputs to be finite"
);
}
}